deerflow-code/offline-backend-20260512/backend/packages/harness/deerflow/knowledge/repository.py
2026-09-07 18:24:55 +08:00

782 lines
35 KiB
Python

"""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