"""Database read/write for the knowledge base. This layer never calls an LLM and never touches the Markdown vault — it only persists/queries the relational mirror (notes, sources, tags). The vault and extraction live in :mod:`deerflow.knowledge.markdown_vault` and :mod:`deerflow.knowledge.extractor`; orchestration lives in the service. """ from __future__ import annotations import uuid from datetime import UTC, datetime from typing import Any from sqlalchemy import delete, func, or_, select from sqlalchemy.ext.asyncio import AsyncSession, async_sessionmaker from deerflow.knowledge.models import ( KnowledgeEmbeddingRow, KnowledgeEntityRow, KnowledgeExtractTemplateRow, KnowledgeFolderRow, KnowledgeGraphBlacklistRow, KnowledgeNoteRow, KnowledgeNoteTagRow, KnowledgeNoteVersionRow, KnowledgeRelationRow, KnowledgeSourceRow, KnowledgeTagRow, ) from deerflow.knowledge.schemas import SourceDraft def _new_id(prefix: str) -> str: return f"{prefix}_{uuid.uuid4().hex[:24]}" def _iso(value: Any) -> Any: if isinstance(value, datetime): return value.isoformat() return value class KnowledgeRepository: """SQLAlchemy-backed repository for knowledge notes/sources/tags.""" def __init__(self, session_factory: async_sessionmaker[AsyncSession]) -> None: self._sf = session_factory # -- serialization ---------------------------------------------------- @staticmethod def _note_to_dict(row: KnowledgeNoteRow) -> dict[str, Any]: d = row.to_dict() d["created_at"] = _iso(d.get("created_at")) d["updated_at"] = _iso(d.get("updated_at")) return d @staticmethod def _source_to_dict(row: KnowledgeSourceRow) -> dict[str, Any]: d = row.to_dict() d["created_at"] = _iso(d.get("created_at")) return d # -- tags ------------------------------------------------------------- async def _get_or_create_tags(self, session: AsyncSession, names: list[str]) -> list[KnowledgeTagRow]: cleaned = [] seen: set[str] = set() for raw in names or []: name = (raw or "").strip() if not name or name.lower() in seen: continue seen.add(name.lower()) cleaned.append(name) if not cleaned: return [] existing = { r.name.lower(): r for r in (await session.execute(select(KnowledgeTagRow).where(func.lower(KnowledgeTagRow.name).in_([n.lower() for n in cleaned])))).scalars() } result: list[KnowledgeTagRow] = [] for name in cleaned: row = existing.get(name.lower()) if row is None: row = KnowledgeTagRow(id=_new_id("kbtag"), name=name, created_at=datetime.now(UTC)) session.add(row) await session.flush() existing[name.lower()] = row result.append(row) return result async def _tags_for(self, session: AsyncSession, note_id: str) -> list[str]: stmt = ( select(KnowledgeTagRow.name) .join(KnowledgeNoteTagRow, KnowledgeNoteTagRow.tag_id == KnowledgeTagRow.id) .where(KnowledgeNoteTagRow.note_id == note_id) .order_by(KnowledgeTagRow.name) ) return list((await session.execute(stmt)).scalars()) async def _tags_for_many(self, session: AsyncSession, note_ids: list[str]) -> dict[str, list[str]]: if not note_ids: return {} stmt = ( select(KnowledgeNoteTagRow.note_id, KnowledgeTagRow.name) .join(KnowledgeTagRow, KnowledgeNoteTagRow.tag_id == KnowledgeTagRow.id) .where(KnowledgeNoteTagRow.note_id.in_(note_ids)) .order_by(KnowledgeTagRow.name) ) out: dict[str, list[str]] = {nid: [] for nid in note_ids} for note_id, name in (await session.execute(stmt)).all(): out.setdefault(note_id, []).append(name) return out async def _replace_note_tags(self, session: AsyncSession, note_id: str, names: list[str]) -> None: await session.execute(delete(KnowledgeNoteTagRow).where(KnowledgeNoteTagRow.note_id == note_id)) tags = await self._get_or_create_tags(session, names) for tag in tags: session.add(KnowledgeNoteTagRow(note_id=note_id, tag_id=tag.id)) # -- notes ------------------------------------------------------------ async def create_note( self, *, title: str, summary: str | None, content_md: str, source_type: str, source_id: str | None, status: str, confidence: float = 0.0, vault_path: str | None = None, folder: str | None = None, tags: list[str] | None = None, created_by: str | None = None, ) -> dict[str, Any]: note_id = _new_id("kb") now = datetime.now(UTC) row = KnowledgeNoteRow( id=note_id, title=title[:512], summary=summary, content_md=content_md, vault_path=vault_path, folder=(folder or None), source_type=source_type, source_id=source_id, status=status, confidence=confidence, created_by=created_by, updated_by=created_by, created_at=now, updated_at=now, ) async with self._sf() as session: session.add(row) await session.flush() await self._replace_note_tags(session, note_id, tags or []) await session.commit() await session.refresh(row) d = self._note_to_dict(row) d["tags"] = await self._tags_for(session, note_id) return d async def set_vault_path(self, note_id: str, vault_path: str) -> None: async with self._sf() as session: row = await session.get(KnowledgeNoteRow, note_id) if row is None: return row.vault_path = vault_path await session.commit() async def update_note( self, note_id: str, *, title: str | None = None, summary: str | None = None, content_md: str | None = None, status: str | None = None, tags: list[str] | None = None, vault_path: str | None = None, folder: str | None = None, set_folder: bool = False, updated_by: str | None = None, ) -> dict[str, Any] | None: async with self._sf() as session: row = await session.get(KnowledgeNoteRow, note_id) if row is None: return None if title is not None: row.title = title[:512] if summary is not None: row.summary = summary if content_md is not None: row.content_md = content_md if status is not None: row.status = status if vault_path is not None: row.vault_path = vault_path # ``set_folder`` lets the caller assign *or clear* (folder=None) the # directory; a bare ``folder=None`` without the flag leaves it as-is. if set_folder: row.folder = (folder or None) if updated_by is not None: row.updated_by = updated_by row.updated_at = datetime.now(UTC) if tags is not None: await self._replace_note_tags(session, note_id, tags) await session.commit() await session.refresh(row) d = self._note_to_dict(row) d["tags"] = await self._tags_for(session, note_id) return d async def set_status(self, note_id: str, status: str, *, updated_by: str | None = None) -> dict[str, Any] | None: return await self.update_note(note_id, status=status, updated_by=updated_by) async def get_note(self, note_id: str, *, with_sources: bool = True) -> dict[str, Any] | None: async with self._sf() as session: row = await session.get(KnowledgeNoteRow, note_id) if row is None: return None d = self._note_to_dict(row) d["tags"] = await self._tags_for(session, note_id) if with_sources: stmt = select(KnowledgeSourceRow).where(KnowledgeSourceRow.note_id == note_id).order_by(KnowledgeSourceRow.created_at) d["sources"] = [self._source_to_dict(r) for r in (await session.execute(stmt)).scalars()] return d async def find_by_vault_path(self, vault_path: str) -> dict[str, Any] | None: """Return the first (non-archived) note whose vault_path matches. ``vault_path`` is matched both with and without the ``.md`` suffix, since wikilink targets are stored without it while ``vault_path`` keeps it. """ stripped = vault_path[:-3] if vault_path.endswith(".md") else vault_path candidates = {vault_path, stripped, f"{stripped}.md"} async with self._sf() as session: stmt = ( select(KnowledgeNoteRow) .where(KnowledgeNoteRow.vault_path.in_(list(candidates))) .where(KnowledgeNoteRow.status != "archived") .limit(1) ) row = (await session.execute(stmt)).scalars().first() if row is None: return None d = self._note_to_dict(row) d["tags"] = await self._tags_for(session, row.id) return d async def find_by_title(self, title: str) -> dict[str, Any] | None: """Return the first (non-archived) note with a case-insensitive title match.""" name = (title or "").strip() if not name: return None async with self._sf() as session: stmt = ( select(KnowledgeNoteRow) .where(func.lower(KnowledgeNoteRow.title) == name.lower()) .where(KnowledgeNoteRow.status != "archived") .order_by(KnowledgeNoteRow.updated_at.desc()) .limit(1) ) row = (await session.execute(stmt)).scalars().first() if row is None: return None d = self._note_to_dict(row) d["tags"] = await self._tags_for(session, row.id) return d async def find_by_source(self, source_type: str, source_id: str) -> dict[str, Any] | None: """Return the first (non-archived) note for a (source_type, source_id).""" async with self._sf() as session: stmt = ( select(KnowledgeNoteRow) .where(KnowledgeNoteRow.source_type == source_type, KnowledgeNoteRow.source_id == source_id) .where(KnowledgeNoteRow.status != "archived") .limit(1) ) row = (await session.execute(stmt)).scalars().first() return self._note_to_dict(row) if row else None async def list_notes( self, *, keyword: str | None = None, tag: str | None = None, entity: str | None = None, source_type: str | None = None, status: str | None = None, folder: str | None = None, include_archived: bool = False, limit: int = 20, offset: int = 0, ) -> tuple[list[dict[str, Any]], int]: async with self._sf() as session: stmt = select(KnowledgeNoteRow) count_stmt = select(func.count()).select_from(KnowledgeNoteRow) conditions = [] if status: conditions.append(KnowledgeNoteRow.status == status) elif not include_archived: conditions.append(KnowledgeNoteRow.status != "archived") if source_type: conditions.append(KnowledgeNoteRow.source_type == source_type) if folder is not None: # Exact folder, plus any nested sub-folders (``folder/…``). if folder == "": conditions.append(KnowledgeNoteRow.folder.is_(None)) else: conditions.append(or_(KnowledgeNoteRow.folder == folder, KnowledgeNoteRow.folder.like(f"{folder}/%"))) if keyword: like = f"%{keyword}%" conditions.append(or_(KnowledgeNoteRow.title.like(like), KnowledgeNoteRow.summary.like(like), KnowledgeNoteRow.content_md.like(like))) if tag: tag_subq = ( select(KnowledgeNoteTagRow.note_id) .join(KnowledgeTagRow, KnowledgeNoteTagRow.tag_id == KnowledgeTagRow.id) .where(func.lower(KnowledgeTagRow.name) == tag.lower()) ) conditions.append(KnowledgeNoteRow.id.in_(tag_subq)) if entity: # Exact entity match via the note→entity relation graph — precise # and index-backed, unlike a content LIKE scan on large vaults. entity_subq = ( select(KnowledgeRelationRow.from_note_id) .join(KnowledgeEntityRow, KnowledgeRelationRow.to_entity_id == KnowledgeEntityRow.id) .where(func.lower(KnowledgeEntityRow.name) == entity.lower()) .where(KnowledgeRelationRow.from_note_id.is_not(None)) ) conditions.append(KnowledgeNoteRow.id.in_(entity_subq)) for cond in conditions: stmt = stmt.where(cond) count_stmt = count_stmt.where(cond) total = (await session.execute(count_stmt)).scalar_one() stmt = stmt.order_by(KnowledgeNoteRow.updated_at.desc()).limit(limit).offset(offset) rows = list((await session.execute(stmt)).scalars()) ids = [r.id for r in rows] tags_map = await self._tags_for_many(session, ids) items = [] for r in rows: d = self._note_to_dict(r) d["tags"] = tags_map.get(r.id, []) items.append(d) return items, int(total) async def list_all(self, *, include_archived: bool = False) -> list[dict[str, Any]]: """Return every note (light fields + tags) for graph export.""" async with self._sf() as session: stmt = select(KnowledgeNoteRow) if not include_archived: stmt = stmt.where(KnowledgeNoteRow.status != "archived") rows = list((await session.execute(stmt)).scalars()) tags_map = await self._tags_for_many(session, [r.id for r in rows]) out = [] for r in rows: d = self._note_to_dict(r) d["tags"] = tags_map.get(r.id, []) out.append(d) return out # -- embeddings (phase 2) -------------------------------------------- async def replace_embeddings(self, note_id: str, chunks: list[tuple[int, str, list[float]]]) -> None: """Replace all embedding rows for a note with ``(index, content, vector)``.""" async with self._sf() as session: await session.execute(delete(KnowledgeEmbeddingRow).where(KnowledgeEmbeddingRow.note_id == note_id)) for idx, content, vector in chunks: session.add( KnowledgeEmbeddingRow( id=_new_id("kbemb"), note_id=note_id, chunk_index=idx, content=content, vector=list(vector), created_at=datetime.now(UTC), ) ) await session.commit() async def delete_embeddings(self, note_id: str) -> None: async with self._sf() as session: await session.execute(delete(KnowledgeEmbeddingRow).where(KnowledgeEmbeddingRow.note_id == note_id)) await session.commit() async def all_embeddings(self) -> list[dict[str, Any]]: """Return every embedding chunk joined to its (non-archived) note.""" async with self._sf() as session: stmt = ( select(KnowledgeEmbeddingRow, KnowledgeNoteRow.title, KnowledgeNoteRow.status) .join(KnowledgeNoteRow, KnowledgeNoteRow.id == KnowledgeEmbeddingRow.note_id) .where(KnowledgeNoteRow.status != "archived") ) out = [] for emb, title, _status in (await session.execute(stmt)).all(): out.append({"note_id": emb.note_id, "chunk_index": emb.chunk_index, "content": emb.content, "vector": emb.vector or [], "title": title}) return out async def has_embeddings(self) -> bool: async with self._sf() as session: row = (await session.execute(select(KnowledgeEmbeddingRow.id).limit(1))).first() return row is not None # -- versions (phase 3) ---------------------------------------------- async def add_version(self, note: dict[str, Any], *, edited_by: str | None) -> None: async with self._sf() as session: session.add( KnowledgeNoteVersionRow( id=_new_id("kbver"), note_id=note["id"], title=note.get("title") or "", summary=note.get("summary"), content_md=note.get("content_md") or "", status=note.get("status") or "approved", tags_json=note.get("tags") or [], edited_by=edited_by, created_at=datetime.now(UTC), ) ) await session.commit() async def list_versions(self, note_id: str) -> list[dict[str, Any]]: async with self._sf() as session: stmt = select(KnowledgeNoteVersionRow).where(KnowledgeNoteVersionRow.note_id == note_id).order_by(KnowledgeNoteVersionRow.created_at.desc()) out = [] for r in (await session.execute(stmt)).scalars(): d = r.to_dict() d["tags"] = d.pop("tags_json", []) d["created_at"] = _iso(d.get("created_at")) out.append(d) return out # -- entities & relations (phase 4) ---------------------------------- async def get_or_create_entity(self, name: str, entity_type: str | None) -> dict[str, Any]: name = (name or "").strip() async with self._sf() as session: stmt = select(KnowledgeEntityRow).where(func.lower(KnowledgeEntityRow.name) == name.lower()) row = (await session.execute(stmt)).scalars().first() if row is None: row = KnowledgeEntityRow(id=_new_id("kbent"), name=name, entity_type=entity_type, created_at=datetime.now(UTC)) session.add(row) await session.commit() await session.refresh(row) return {"id": row.id, "name": row.name, "entity_type": row.entity_type} async def add_relation( self, *, from_note_id: str | None = None, to_note_id: str | None = None, from_entity_id: str | None = None, to_entity_id: str | None = None, relation_type: str, weight: float = 1.0, ) -> None: async with self._sf() as session: session.add( KnowledgeRelationRow( id=_new_id("kbrel"), from_note_id=from_note_id, to_note_id=to_note_id, from_entity_id=from_entity_id, to_entity_id=to_entity_id, relation_type=relation_type, weight=weight, created_at=datetime.now(UTC), ) ) await session.commit() async def clear_note_relations(self, note_id: str) -> None: async with self._sf() as session: await session.execute(delete(KnowledgeRelationRow).where(KnowledgeRelationRow.from_note_id == note_id)) await session.commit() async def list_note_entities(self, note_id: str) -> list[dict[str, Any]]: """Return the entities mentioned by a note (for note-detail display).""" async with self._sf() as session: stmt = ( select(KnowledgeEntityRow.name, KnowledgeEntityRow.entity_type) .join(KnowledgeRelationRow, KnowledgeRelationRow.to_entity_id == KnowledgeEntityRow.id) .where(KnowledgeRelationRow.from_note_id == note_id) .order_by(KnowledgeEntityRow.name) ) seen: set[str] = set() out: list[dict[str, Any]] = [] for name, etype in (await session.execute(stmt)).all(): key = (name or "").lower() if not name or key in seen: continue seen.add(key) out.append({"name": name, "entity_type": etype}) return out async def prune_orphan_entities(self) -> int: """Delete entity rows no longer referenced by any relation. Returns count.""" async with self._sf() as session: referenced: set[str] = set() for col in (KnowledgeRelationRow.from_entity_id, KnowledgeRelationRow.to_entity_id): for (eid,) in (await session.execute(select(col).where(col.isnot(None)))).all(): referenced.add(eid) all_ids = list((await session.execute(select(KnowledgeEntityRow.id))).scalars()) orphans = [eid for eid in all_ids if eid not in referenced] if orphans: await session.execute(delete(KnowledgeEntityRow).where(KnowledgeEntityRow.id.in_(orphans))) await session.commit() return len(orphans) async def list_entities(self) -> list[dict[str, Any]]: async with self._sf() as session: rows = (await session.execute(select(KnowledgeEntityRow))).scalars() return [{"id": r.id, "name": r.name, "entity_type": r.entity_type} for r in rows] # -- graph blacklist --------------------------------------------------- async def list_graph_blacklist(self) -> list[dict[str, Any]]: """Return every blacklisted graph node label, newest first.""" async with self._sf() as session: stmt = select(KnowledgeGraphBlacklistRow).order_by(KnowledgeGraphBlacklistRow.created_at.desc()) rows = (await session.execute(stmt)).scalars() return [{"id": r.id, "label": r.label, "created_by": r.created_by, "created_at": _iso(r.created_at)} for r in rows] async def add_graph_blacklist(self, labels: list[str], *, created_by: str | None = None) -> int: """Add labels to the graph blacklist (trimmed, case-insensitive dedup). Returns the count actually added.""" cleaned: list[str] = [] seen: set[str] = set() for raw in labels or []: label = (raw or "").strip() if not label or label.lower() in seen: continue seen.add(label.lower()) cleaned.append(label) if not cleaned: return 0 async with self._sf() as session: existing = { value for (value,) in ( await session.execute(select(func.lower(KnowledgeGraphBlacklistRow.label)).where(func.lower(KnowledgeGraphBlacklistRow.label).in_([label.lower() for label in cleaned]))) ).all() } added = 0 for label in cleaned: if label.lower() in existing: continue session.add(KnowledgeGraphBlacklistRow(id=_new_id("kbblk"), label=label, created_by=created_by, created_at=datetime.now(UTC))) added += 1 if added: await session.commit() return added async def remove_graph_blacklist(self, entry_id: str) -> bool: async with self._sf() as session: result = await session.execute(delete(KnowledgeGraphBlacklistRow).where(KnowledgeGraphBlacklistRow.id == entry_id)) await session.commit() return bool(result.rowcount) async def list_relations(self) -> list[dict[str, Any]]: async with self._sf() as session: rows = (await session.execute(select(KnowledgeRelationRow))).scalars() return [ { "id": r.id, "from_note_id": r.from_note_id, "to_note_id": r.to_note_id, "from_entity_id": r.from_entity_id, "to_entity_id": r.to_entity_id, "relation_type": r.relation_type, "weight": r.weight, } for r in rows ] # -- sources ---------------------------------------------------------- async def add_sources(self, note_id: str, sources: list[SourceDraft], *, created_by: str | None = None) -> None: if not sources: return import json async with self._sf() as session: for src in sources: session.add( KnowledgeSourceRow( id=_new_id("kbsrc"), note_id=note_id, source_type=src.source_type, thread_id=src.thread_id, message_id=src.message_id, tool_name=src.tool_name, title=(src.title or "")[:512] or None, url=src.url, snippet=src.snippet, raw_json=json.dumps(src.raw, ensure_ascii=False) if src.raw else None, created_by=created_by, created_at=datetime.now(UTC), ) ) await session.commit() # -- folders ---------------------------------------------------------- async def list_folders(self) -> list[dict[str, Any]]: """Return every managed directory, ordered by sort_order then path.""" async with self._sf() as session: stmt = select(KnowledgeFolderRow).order_by(KnowledgeFolderRow.sort_order, KnowledgeFolderRow.path) return [{"id": r.id, "path": r.path, "name": r.name, "sort_order": r.sort_order} for r in (await session.execute(stmt)).scalars()] async def create_folder(self, path: str, *, created_by: str | None = None) -> dict[str, Any]: """Create a directory by its full slash-joined path (idempotent on path). Parent directories are auto-created so ``投研/行业/寿险`` materializes all three levels. Returns the leaf folder row. """ segments = [s.strip() for s in (path or "").split("/") if s.strip()] if not segments: raise ValueError("folder path is empty") async with self._sf() as session: acc: list[str] = [] leaf: KnowledgeFolderRow | None = None for i, seg in enumerate(segments): acc.append(seg) full = "/".join(acc) row = await self._folder_by_path(session, full) if row is None: row = KnowledgeFolderRow(id=_new_id("kbdir"), path=full, name=seg, sort_order=i, created_by=created_by, created_at=datetime.now(UTC)) session.add(row) await session.flush() leaf = row # Capture before commit: the async session expires attributes on # commit, and re-reading them outside an awaited load would error. result = {"id": leaf.id, "path": leaf.path, "name": leaf.name, "sort_order": leaf.sort_order} await session.commit() return result @staticmethod async def _folder_by_path(session: AsyncSession, path: str) -> KnowledgeFolderRow | None: return (await session.execute(select(KnowledgeFolderRow).where(KnowledgeFolderRow.path == path))).scalars().first() async def rename_folder(self, folder_id: str, new_path: str) -> dict[str, Any] | None: """Rename/move a folder and re-point every descendant folder + note. Both the folder rows under ``old/…`` and notes whose ``folder`` is ``old`` or ``old/…`` are rewritten to the new prefix, so the move is atomic. """ new_path = "/".join(s.strip() for s in (new_path or "").split("/") if s.strip()) if not new_path: raise ValueError("folder path is empty") async with self._sf() as session: row = await session.get(KnowledgeFolderRow, folder_id) if row is None: return None old = row.path if new_path == old: return {"id": row.id, "path": row.path, "name": row.name, "sort_order": row.sort_order} for f in (await session.execute(select(KnowledgeFolderRow).where(or_(KnowledgeFolderRow.path == old, KnowledgeFolderRow.path.like(f"{old}/%"))))).scalars(): f.path = new_path + f.path[len(old):] f.name = f.path.split("/")[-1] for n in (await session.execute(select(KnowledgeNoteRow).where(or_(KnowledgeNoteRow.folder == old, KnowledgeNoteRow.folder.like(f"{old}/%"))))).scalars(): n.folder = new_path + (n.folder[len(old):] if n.folder else "") await session.commit() refreshed = await session.get(KnowledgeFolderRow, folder_id) return {"id": refreshed.id, "path": refreshed.path, "name": refreshed.name, "sort_order": refreshed.sort_order} async def delete_folder(self, folder_id: str, *, reassign_to: str | None = None) -> bool: """Delete a folder (and its sub-folders); notes inside move to ``reassign_to`` (or unfiled).""" async with self._sf() as session: row = await session.get(KnowledgeFolderRow, folder_id) if row is None: return False old = row.path target = (reassign_to or "").strip() or None for n in (await session.execute(select(KnowledgeNoteRow).where(or_(KnowledgeNoteRow.folder == old, KnowledgeNoteRow.folder.like(f"{old}/%"))))).scalars(): n.folder = target await session.execute(delete(KnowledgeFolderRow).where(or_(KnowledgeFolderRow.path == old, KnowledgeFolderRow.path.like(f"{old}/%")))) await session.commit() return True # -- extraction templates --------------------------------------------- @staticmethod def _template_to_dict(row: KnowledgeExtractTemplateRow) -> dict[str, Any]: d = row.to_dict() d["created_at"] = _iso(d.get("created_at")) d["updated_at"] = _iso(d.get("updated_at")) return d async def list_templates(self, *, scope: str | None = None, enabled_only: bool = False) -> list[dict[str, Any]]: async with self._sf() as session: stmt = select(KnowledgeExtractTemplateRow) if scope: stmt = stmt.where(or_(KnowledgeExtractTemplateRow.scope == scope, KnowledgeExtractTemplateRow.scope == "both")) if enabled_only: stmt = stmt.where(KnowledgeExtractTemplateRow.enabled.is_(True)) stmt = stmt.order_by(KnowledgeExtractTemplateRow.is_default.desc(), KnowledgeExtractTemplateRow.created_at) return [self._template_to_dict(r) for r in (await session.execute(stmt)).scalars()] async def get_template(self, template_id: str) -> dict[str, Any] | None: async with self._sf() as session: row = await session.get(KnowledgeExtractTemplateRow, template_id) return self._template_to_dict(row) if row else None async def get_default_template(self, scope: str) -> dict[str, Any] | None: """Return the default enabled template applicable to ``scope`` (thread/document).""" async with self._sf() as session: stmt = ( select(KnowledgeExtractTemplateRow) .where(KnowledgeExtractTemplateRow.enabled.is_(True)) .where(KnowledgeExtractTemplateRow.is_default.is_(True)) .where(or_(KnowledgeExtractTemplateRow.scope == scope, KnowledgeExtractTemplateRow.scope == "both")) .limit(1) ) row = (await session.execute(stmt)).scalars().first() return self._template_to_dict(row) if row else None async def template_name_taken(self, name: str, *, exclude_id: str | None = None) -> bool: async with self._sf() as session: stmt = select(KnowledgeExtractTemplateRow.id).where(func.lower(KnowledgeExtractTemplateRow.name) == (name or "").strip().lower()) if exclude_id: stmt = stmt.where(KnowledgeExtractTemplateRow.id != exclude_id) return (await session.execute(stmt)).first() is not None async def create_template(self, *, name: str, system_prompt: str, scope: str = "both", description: str | None = None, is_default: bool = False, enabled: bool = True, created_by: str | None = None) -> dict[str, Any]: async with self._sf() as session: if is_default: await self._clear_default(session, scope) row = KnowledgeExtractTemplateRow( id=_new_id("kbtpl"), name=name.strip()[:256], description=description, scope=scope, system_prompt=system_prompt, is_default=is_default, enabled=enabled, created_by=created_by, updated_by=created_by, created_at=datetime.now(UTC), updated_at=datetime.now(UTC), ) session.add(row) await session.commit() await session.refresh(row) return self._template_to_dict(row) async def update_template(self, template_id: str, *, name: str | None = None, system_prompt: str | None = None, scope: str | None = None, description: str | None = None, is_default: bool | None = None, enabled: bool | None = None, updated_by: str | None = None) -> dict[str, Any] | None: async with self._sf() as session: row = await session.get(KnowledgeExtractTemplateRow, template_id) if row is None: return None if name is not None: row.name = name.strip()[:256] if system_prompt is not None: row.system_prompt = system_prompt if scope is not None: row.scope = scope if description is not None: row.description = description if enabled is not None: row.enabled = enabled if is_default: await self._clear_default(session, scope or row.scope) row.is_default = True elif is_default is False: row.is_default = False row.updated_by = updated_by row.updated_at = datetime.now(UTC) await session.commit() await session.refresh(row) return self._template_to_dict(row) async def delete_template(self, template_id: str) -> bool: async with self._sf() as session: result = await session.execute(delete(KnowledgeExtractTemplateRow).where(KnowledgeExtractTemplateRow.id == template_id)) await session.commit() return bool(result.rowcount) @staticmethod async def _clear_default(session: AsyncSession, scope: str) -> None: """Unset is_default on any template that would collide with a new default for ``scope``.""" scopes = {scope, "both"} if scope != "both" else {"thread", "document", "both"} for row in (await session.execute(select(KnowledgeExtractTemplateRow).where(KnowledgeExtractTemplateRow.is_default.is_(True)).where(KnowledgeExtractTemplateRow.scope.in_(scopes)))).scalars(): row.is_default = False