feat: add sqlite persistence model
This commit is contained in:
@@ -0,0 +1 @@
|
||||
"""Chat persistence and streaming feature."""
|
||||
@@ -0,0 +1,78 @@
|
||||
from datetime import datetime
|
||||
|
||||
from sqlalchemy import (
|
||||
CheckConstraint,
|
||||
DateTime,
|
||||
Float,
|
||||
ForeignKey,
|
||||
Index,
|
||||
Integer,
|
||||
String,
|
||||
Text,
|
||||
UniqueConstraint,
|
||||
)
|
||||
from sqlalchemy.orm import Mapped, mapped_column
|
||||
|
||||
from kbqa.database import Base, utc_now
|
||||
from kbqa.documents.models import new_uuid
|
||||
|
||||
|
||||
class ChatSession(Base):
|
||||
__tablename__ = "chat_sessions"
|
||||
__table_args__ = (Index("ix_chat_sessions_updated_at", "updated_at"),)
|
||||
|
||||
id: Mapped[str] = mapped_column(String(36), primary_key=True, default=new_uuid)
|
||||
title: Mapped[str | None] = mapped_column(Text)
|
||||
created_at: Mapped[datetime] = mapped_column(
|
||||
DateTime(timezone=True), nullable=False, default=utc_now
|
||||
)
|
||||
updated_at: Mapped[datetime] = mapped_column(
|
||||
DateTime(timezone=True), nullable=False, default=utc_now, onupdate=utc_now
|
||||
)
|
||||
|
||||
|
||||
class Message(Base):
|
||||
__tablename__ = "messages"
|
||||
__table_args__ = (
|
||||
CheckConstraint("role IN ('user','assistant')", name="ck_messages_role"),
|
||||
CheckConstraint("status IN ('generating','completed','failed')", name="ck_messages_status"),
|
||||
UniqueConstraint("session_id", "order_index", name="uq_messages_session_order"),
|
||||
Index("ix_messages_session_created_at", "session_id", "created_at"),
|
||||
)
|
||||
|
||||
id: Mapped[str] = mapped_column(String(36), primary_key=True, default=new_uuid)
|
||||
session_id: Mapped[str] = mapped_column(
|
||||
String(36), ForeignKey("chat_sessions.id", ondelete="CASCADE"), nullable=False
|
||||
)
|
||||
role: Mapped[str] = mapped_column(String(16), nullable=False)
|
||||
content: Mapped[str] = mapped_column(Text, nullable=False, default="")
|
||||
status: Mapped[str] = mapped_column(String(16), nullable=False)
|
||||
error_code: Mapped[str | None] = mapped_column(String(64))
|
||||
order_index: Mapped[int] = mapped_column(Integer, nullable=False)
|
||||
created_at: Mapped[datetime] = mapped_column(
|
||||
DateTime(timezone=True), nullable=False, default=utc_now
|
||||
)
|
||||
|
||||
|
||||
class MessageSource(Base):
|
||||
__tablename__ = "message_sources"
|
||||
__table_args__ = (
|
||||
UniqueConstraint("message_id", "chunk_id", name="uq_message_sources_message_chunk"),
|
||||
Index("ix_message_sources_message_id", "message_id"),
|
||||
)
|
||||
|
||||
id: Mapped[str] = mapped_column(String(36), primary_key=True, default=new_uuid)
|
||||
message_id: Mapped[str] = mapped_column(
|
||||
String(36), ForeignKey("messages.id", ondelete="CASCADE"), nullable=False
|
||||
)
|
||||
document_id: Mapped[str] = mapped_column(
|
||||
String(36), ForeignKey("documents.id", ondelete="CASCADE"), nullable=False
|
||||
)
|
||||
chunk_id: Mapped[str] = mapped_column(
|
||||
String(36), ForeignKey("chunks.id", ondelete="CASCADE"), nullable=False
|
||||
)
|
||||
citation_label: Mapped[str] = mapped_column(String(8), nullable=False)
|
||||
score: Mapped[float] = mapped_column(Float, nullable=False)
|
||||
created_at: Mapped[datetime] = mapped_column(
|
||||
DateTime(timezone=True), nullable=False, default=utc_now
|
||||
)
|
||||
@@ -0,0 +1,164 @@
|
||||
import re
|
||||
from dataclasses import dataclass
|
||||
|
||||
from sqlalchemy import delete, func, select
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from kbqa.chat.models import ChatSession, Message, MessageSource
|
||||
from kbqa.database import utc_now
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class MessageSourceInput:
|
||||
document_id: str
|
||||
chunk_id: str
|
||||
citation_label: str
|
||||
score: float
|
||||
|
||||
|
||||
class ChatRepository:
|
||||
def __init__(self, session: AsyncSession) -> None:
|
||||
self.session = session
|
||||
|
||||
async def create_session(self, title: str | None) -> ChatSession:
|
||||
session = ChatSession(title=title.strip() if title and title.strip() else None)
|
||||
self.session.add(session)
|
||||
await self.session.commit()
|
||||
await self.session.refresh(session)
|
||||
return session
|
||||
|
||||
async def get_session(self, session_id: str) -> ChatSession | None:
|
||||
return await self.session.get(ChatSession, session_id)
|
||||
|
||||
async def page_sessions(self, offset: int, limit: int) -> tuple[list[ChatSession], int]:
|
||||
total = await self.session.scalar(select(func.count(ChatSession.id)))
|
||||
rows = await self.session.scalars(
|
||||
select(ChatSession).order_by(ChatSession.updated_at.desc()).offset(offset).limit(limit)
|
||||
)
|
||||
return list(rows), int(total or 0)
|
||||
|
||||
async def load_recent_messages(self, session_id: str, limit: int) -> list[Message]:
|
||||
rows = list(
|
||||
await self.session.scalars(
|
||||
select(Message)
|
||||
.where(Message.session_id == session_id, Message.status == "completed")
|
||||
.order_by(Message.order_index.desc())
|
||||
.limit(limit)
|
||||
)
|
||||
)
|
||||
rows.reverse()
|
||||
return rows
|
||||
|
||||
async def page_messages(
|
||||
self,
|
||||
session_id: str,
|
||||
offset: int,
|
||||
limit: int,
|
||||
) -> tuple[list[Message], int]:
|
||||
total = await self.session.scalar(
|
||||
select(func.count(Message.id)).where(Message.session_id == session_id)
|
||||
)
|
||||
rows = await self.session.scalars(
|
||||
select(Message)
|
||||
.where(Message.session_id == session_id)
|
||||
.order_by(Message.order_index)
|
||||
.offset(offset)
|
||||
.limit(limit)
|
||||
)
|
||||
return list(rows), int(total or 0)
|
||||
|
||||
async def sources_for_messages(self, message_ids: list[str]) -> dict[str, list[MessageSource]]:
|
||||
if not message_ids:
|
||||
return {}
|
||||
sources = await self.session.scalars(
|
||||
select(MessageSource)
|
||||
.where(MessageSource.message_id.in_(message_ids))
|
||||
.order_by(MessageSource.citation_label)
|
||||
)
|
||||
result: dict[str, list[MessageSource]] = {}
|
||||
for source in sources:
|
||||
result.setdefault(source.message_id, []).append(source)
|
||||
return result
|
||||
|
||||
async def begin_turn(self, session_id: str, query: str) -> tuple[Message, Message]:
|
||||
chat_session = await self.session.get(ChatSession, session_id)
|
||||
if chat_session is None:
|
||||
raise LookupError("session not found")
|
||||
maximum = await self.session.scalar(
|
||||
select(func.max(Message.order_index)).where(Message.session_id == session_id)
|
||||
)
|
||||
next_order = int(maximum if maximum is not None else -1) + 1
|
||||
user = Message(
|
||||
session_id=session_id,
|
||||
role="user",
|
||||
content=query,
|
||||
status="completed",
|
||||
order_index=next_order,
|
||||
)
|
||||
assistant = Message(
|
||||
session_id=session_id,
|
||||
role="assistant",
|
||||
content="",
|
||||
status="generating",
|
||||
order_index=next_order + 1,
|
||||
)
|
||||
if chat_session.title is None:
|
||||
chat_session.title = query.strip()[:30]
|
||||
chat_session.updated_at = utc_now()
|
||||
self.session.add_all([user, assistant])
|
||||
await self.session.commit()
|
||||
await self.session.refresh(user)
|
||||
await self.session.refresh(assistant)
|
||||
return user, assistant
|
||||
|
||||
async def complete_answer(
|
||||
self,
|
||||
message_id: str,
|
||||
content: str,
|
||||
sources: list[MessageSourceInput],
|
||||
) -> None:
|
||||
message = await self.session.get(Message, message_id)
|
||||
if message is None:
|
||||
raise LookupError("message not found")
|
||||
await self.session.execute(
|
||||
delete(MessageSource).where(MessageSource.message_id == message_id)
|
||||
)
|
||||
message.content = content
|
||||
message.status = "completed"
|
||||
message.error_code = None
|
||||
self.session.add_all(
|
||||
[
|
||||
MessageSource(
|
||||
message_id=message_id,
|
||||
document_id=source.document_id,
|
||||
chunk_id=source.chunk_id,
|
||||
citation_label=source.citation_label,
|
||||
score=source.score,
|
||||
)
|
||||
for source in sources
|
||||
]
|
||||
)
|
||||
chat_session = await self.session.get(ChatSession, message.session_id)
|
||||
if chat_session is not None:
|
||||
chat_session.updated_at = utc_now()
|
||||
await self.session.commit()
|
||||
|
||||
async def fail_answer(self, message_id: str, error_code: str) -> None:
|
||||
message = await self.session.get(Message, message_id)
|
||||
if message is None:
|
||||
return
|
||||
message.status = "failed"
|
||||
message.error_code = error_code
|
||||
await self.session.commit()
|
||||
|
||||
async def delete_session(self, session_id: str) -> bool:
|
||||
chat_session = await self.session.get(ChatSession, session_id)
|
||||
if chat_session is None:
|
||||
return False
|
||||
await self.session.delete(chat_session)
|
||||
await self.session.commit()
|
||||
return True
|
||||
|
||||
@staticmethod
|
||||
def citation_labels(content: str) -> set[str]:
|
||||
return set(re.findall(r"\[(S\d+)]", content))
|
||||
@@ -1,4 +1,5 @@
|
||||
from collections.abc import AsyncIterator
|
||||
from datetime import UTC, datetime
|
||||
|
||||
from fastapi import Request
|
||||
from sqlalchemy import event
|
||||
@@ -16,6 +17,10 @@ class Base(DeclarativeBase):
|
||||
pass
|
||||
|
||||
|
||||
def utc_now() -> datetime:
|
||||
return datetime.now(UTC)
|
||||
|
||||
|
||||
class Database:
|
||||
def __init__(self, url: str) -> None:
|
||||
self.engine: AsyncEngine = create_async_engine(url, pool_pre_ping=True)
|
||||
|
||||
@@ -0,0 +1 @@
|
||||
"""Document management feature."""
|
||||
@@ -0,0 +1,92 @@
|
||||
from datetime import datetime
|
||||
from uuid import uuid4
|
||||
|
||||
from sqlalchemy import (
|
||||
CheckConstraint,
|
||||
DateTime,
|
||||
ForeignKey,
|
||||
Index,
|
||||
Integer,
|
||||
String,
|
||||
Text,
|
||||
UniqueConstraint,
|
||||
)
|
||||
from sqlalchemy.orm import Mapped, mapped_column
|
||||
|
||||
from kbqa.database import Base, utc_now
|
||||
|
||||
|
||||
def new_uuid() -> str:
|
||||
return str(uuid4())
|
||||
|
||||
|
||||
class Document(Base):
|
||||
__tablename__ = "documents"
|
||||
__table_args__ = (
|
||||
CheckConstraint(
|
||||
"status IN ('pending','indexing','indexed','failed','deleting')",
|
||||
name="ck_documents_status",
|
||||
),
|
||||
Index("ix_documents_status_updated_at", "status", "updated_at"),
|
||||
)
|
||||
|
||||
id: Mapped[str] = mapped_column(String(36), primary_key=True, default=new_uuid)
|
||||
filename: Mapped[str] = mapped_column(Text, nullable=False)
|
||||
stored_path: Mapped[str] = mapped_column(Text, nullable=False)
|
||||
file_type: Mapped[str] = mapped_column(String(16), nullable=False)
|
||||
mime_type: Mapped[str] = mapped_column(String(128), nullable=False)
|
||||
file_size: Mapped[int] = mapped_column(Integer, nullable=False)
|
||||
sha256: Mapped[str] = mapped_column(String(64), nullable=False)
|
||||
status: Mapped[str] = mapped_column(String(16), nullable=False, default="pending")
|
||||
chunk_count: Mapped[int] = mapped_column(Integer, nullable=False, default=0)
|
||||
error_code: Mapped[str | None] = mapped_column(String(64))
|
||||
error_message: Mapped[str | None] = mapped_column(Text)
|
||||
created_at: Mapped[datetime] = mapped_column(
|
||||
DateTime(timezone=True), nullable=False, default=utc_now
|
||||
)
|
||||
updated_at: Mapped[datetime] = mapped_column(
|
||||
DateTime(timezone=True), nullable=False, default=utc_now, onupdate=utc_now
|
||||
)
|
||||
|
||||
|
||||
class IndexJob(Base):
|
||||
__tablename__ = "index_jobs"
|
||||
__table_args__ = (
|
||||
CheckConstraint(
|
||||
"status IN ('queued','running','succeeded','failed')",
|
||||
name="ck_index_jobs_status",
|
||||
),
|
||||
)
|
||||
|
||||
id: Mapped[str] = mapped_column(String(36), primary_key=True, default=new_uuid)
|
||||
document_id: Mapped[str] = mapped_column(
|
||||
String(36), ForeignKey("documents.id", ondelete="CASCADE"), nullable=False, unique=True
|
||||
)
|
||||
status: Mapped[str] = mapped_column(String(16), nullable=False, default="queued")
|
||||
attempt_count: Mapped[int] = mapped_column(Integer, nullable=False, default=0)
|
||||
last_error: Mapped[str | None] = mapped_column(Text)
|
||||
created_at: Mapped[datetime] = mapped_column(
|
||||
DateTime(timezone=True), nullable=False, default=utc_now
|
||||
)
|
||||
started_at: Mapped[datetime | None] = mapped_column(DateTime(timezone=True))
|
||||
finished_at: Mapped[datetime | None] = mapped_column(DateTime(timezone=True))
|
||||
|
||||
|
||||
class Chunk(Base):
|
||||
__tablename__ = "chunks"
|
||||
__table_args__ = (
|
||||
UniqueConstraint("document_id", "order_index", name="uq_chunks_document_order"),
|
||||
Index("ix_chunks_document_id", "document_id"),
|
||||
)
|
||||
|
||||
id: Mapped[str] = mapped_column(String(36), primary_key=True)
|
||||
document_id: Mapped[str] = mapped_column(
|
||||
String(36), ForeignKey("documents.id", ondelete="CASCADE"), nullable=False
|
||||
)
|
||||
order_index: Mapped[int] = mapped_column(Integer, nullable=False)
|
||||
content: Mapped[str] = mapped_column(Text, nullable=False)
|
||||
page_number: Mapped[int | None] = mapped_column(Integer)
|
||||
char_count: Mapped[int] = mapped_column(Integer, nullable=False)
|
||||
created_at: Mapped[datetime] = mapped_column(
|
||||
DateTime(timezone=True), nullable=False, default=utc_now
|
||||
)
|
||||
@@ -0,0 +1,195 @@
|
||||
from dataclasses import dataclass
|
||||
|
||||
from sqlalchemy import delete, func, select
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from kbqa.database import utc_now
|
||||
from kbqa.documents.models import Chunk, Document, IndexJob, new_uuid
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class StoredUpload:
|
||||
document_id: str
|
||||
filename: str
|
||||
stored_path: str
|
||||
file_type: str
|
||||
mime_type: str
|
||||
file_size: int
|
||||
sha256: str
|
||||
|
||||
|
||||
class DocumentRepository:
|
||||
def __init__(self, session: AsyncSession) -> None:
|
||||
self.session = session
|
||||
|
||||
async def create_pending(self, upload: StoredUpload) -> tuple[Document, IndexJob]:
|
||||
document = Document(
|
||||
id=upload.document_id,
|
||||
filename=upload.filename,
|
||||
stored_path=upload.stored_path,
|
||||
file_type=upload.file_type,
|
||||
mime_type=upload.mime_type,
|
||||
file_size=upload.file_size,
|
||||
sha256=upload.sha256,
|
||||
status="pending",
|
||||
)
|
||||
job = IndexJob(id=new_uuid(), document_id=document.id, status="queued")
|
||||
self.session.add_all([document, job])
|
||||
await self.session.commit()
|
||||
await self.session.refresh(document)
|
||||
await self.session.refresh(job)
|
||||
return document, job
|
||||
|
||||
async def page(
|
||||
self,
|
||||
status: str | None,
|
||||
offset: int,
|
||||
limit: int,
|
||||
) -> tuple[list[Document], int]:
|
||||
filters = [Document.status == status] if status else []
|
||||
total = await self.session.scalar(select(func.count(Document.id)).where(*filters))
|
||||
result = await self.session.scalars(
|
||||
select(Document)
|
||||
.where(*filters)
|
||||
.order_by(Document.created_at.desc())
|
||||
.offset(offset)
|
||||
.limit(limit)
|
||||
)
|
||||
return list(result), int(total or 0)
|
||||
|
||||
async def get(self, document_id: str) -> Document | None:
|
||||
return await self.session.get(Document, document_id)
|
||||
|
||||
async def get_job_for_document(self, document_id: str) -> IndexJob | None:
|
||||
return await self.session.scalar(
|
||||
select(IndexJob).where(IndexJob.document_id == document_id)
|
||||
)
|
||||
|
||||
async def page_chunks(
|
||||
self,
|
||||
document_id: str,
|
||||
offset: int,
|
||||
limit: int,
|
||||
) -> tuple[list[Chunk], int]:
|
||||
total = await self.session.scalar(
|
||||
select(func.count(Chunk.id)).where(Chunk.document_id == document_id)
|
||||
)
|
||||
result = await self.session.scalars(
|
||||
select(Chunk)
|
||||
.where(Chunk.document_id == document_id)
|
||||
.order_by(Chunk.order_index)
|
||||
.offset(offset)
|
||||
.limit(limit)
|
||||
)
|
||||
return list(result), int(total or 0)
|
||||
|
||||
async def claim_job(self, job_id: str) -> IndexJob | None:
|
||||
job = await self.session.get(IndexJob, job_id)
|
||||
if job is None or job.status != "queued":
|
||||
return None
|
||||
document = await self.session.get(Document, job.document_id)
|
||||
if document is None or document.status not in {"pending", "failed"}:
|
||||
return None
|
||||
now = utc_now()
|
||||
job.status = "running"
|
||||
job.attempt_count += 1
|
||||
job.started_at = now
|
||||
job.finished_at = None
|
||||
job.last_error = None
|
||||
document.status = "indexing"
|
||||
document.error_code = None
|
||||
document.error_message = None
|
||||
document.updated_at = now
|
||||
await self.session.commit()
|
||||
return job
|
||||
|
||||
async def replace_chunks(self, document_id: str, chunks: list[Chunk]) -> None:
|
||||
await self.session.execute(delete(Chunk).where(Chunk.document_id == document_id))
|
||||
self.session.add_all(chunks)
|
||||
await self.session.commit()
|
||||
|
||||
async def mark_indexed(self, document_id: str, chunk_count: int) -> None:
|
||||
document = await self.session.get(Document, document_id)
|
||||
job = await self.get_job_for_document(document_id)
|
||||
if document is None or job is None:
|
||||
return
|
||||
now = utc_now()
|
||||
document.status = "indexed"
|
||||
document.chunk_count = chunk_count
|
||||
document.updated_at = now
|
||||
document.error_code = None
|
||||
document.error_message = None
|
||||
job.status = "succeeded"
|
||||
job.finished_at = now
|
||||
job.last_error = None
|
||||
await self.session.commit()
|
||||
|
||||
async def mark_failed(self, document_id: str, code: str, message: str) -> None:
|
||||
document = await self.session.get(Document, document_id)
|
||||
job = await self.get_job_for_document(document_id)
|
||||
if document is None or job is None:
|
||||
return
|
||||
now = utc_now()
|
||||
document.status = "failed"
|
||||
document.error_code = code
|
||||
document.error_message = message
|
||||
document.updated_at = now
|
||||
job.status = "failed"
|
||||
job.finished_at = now
|
||||
job.last_error = message
|
||||
await self.session.commit()
|
||||
|
||||
async def requeue_unfinished(self) -> list[str]:
|
||||
jobs = list(
|
||||
await self.session.scalars(
|
||||
select(IndexJob).where(IndexJob.status.in_(("queued", "running")))
|
||||
)
|
||||
)
|
||||
now = utc_now()
|
||||
for job in jobs:
|
||||
job.status = "queued"
|
||||
job.started_at = None
|
||||
job.finished_at = None
|
||||
document = await self.session.get(Document, job.document_id)
|
||||
if document is not None:
|
||||
document.status = "pending"
|
||||
document.updated_at = now
|
||||
await self.session.commit()
|
||||
return [job.id for job in jobs]
|
||||
|
||||
async def requeue_failed(self, document_id: str) -> IndexJob | None:
|
||||
document = await self.session.get(Document, document_id)
|
||||
job = await self.get_job_for_document(document_id)
|
||||
if document is None or job is None or document.status != "failed":
|
||||
return None
|
||||
document.status = "pending"
|
||||
document.error_code = None
|
||||
document.error_message = None
|
||||
document.updated_at = utc_now()
|
||||
job.status = "queued"
|
||||
job.last_error = None
|
||||
job.finished_at = None
|
||||
await self.session.commit()
|
||||
return job
|
||||
|
||||
async def mark_deleting(self, document_id: str) -> Document | None:
|
||||
document = await self.session.get(Document, document_id)
|
||||
if document is None:
|
||||
return None
|
||||
document.status = "deleting"
|
||||
document.updated_at = utc_now()
|
||||
await self.session.commit()
|
||||
return document
|
||||
|
||||
async def deleting_documents(self) -> list[Document]:
|
||||
return list(
|
||||
await self.session.scalars(select(Document).where(Document.status == "deleting"))
|
||||
)
|
||||
|
||||
async def delete_document(self, document_id: str) -> bool:
|
||||
document = await self.session.get(Document, document_id)
|
||||
if document is None:
|
||||
return False
|
||||
await self.session.delete(document)
|
||||
await self.session.commit()
|
||||
return True
|
||||
@@ -0,0 +1,6 @@
|
||||
"""Import all ORM models so Alembic can discover metadata."""
|
||||
|
||||
from kbqa.chat.models import ChatSession, Message, MessageSource
|
||||
from kbqa.documents.models import Chunk, Document, IndexJob
|
||||
|
||||
__all__ = ["ChatSession", "Chunk", "Document", "IndexJob", "Message", "MessageSource"]
|
||||
Reference in New Issue
Block a user