782 lines
35 KiB
Python
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
|