Persist regenerated answer versions
This commit is contained in:
@@ -6,7 +6,8 @@ import os
|
||||
import sqlite3
|
||||
from uuid import UUID
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
from contextlib import contextmanager
|
||||
from typing import Any, Iterator
|
||||
|
||||
|
||||
def _database_path() -> Path:
|
||||
@@ -27,11 +28,16 @@ LOCAL_USER_ID = str(
|
||||
)
|
||||
|
||||
|
||||
def _connect() -> sqlite3.Connection:
|
||||
@contextmanager
|
||||
def _connect() -> Iterator[sqlite3.Connection]:
|
||||
connection = sqlite3.connect(DATABASE_PATH, timeout=10)
|
||||
connection.row_factory = sqlite3.Row
|
||||
connection.execute("PRAGMA foreign_keys = ON")
|
||||
return connection
|
||||
try:
|
||||
with connection:
|
||||
yield connection
|
||||
finally:
|
||||
connection.close()
|
||||
|
||||
|
||||
def initialize_database() -> None:
|
||||
@@ -70,13 +76,28 @@ def initialize_database() -> None:
|
||||
conversation_id TEXT NOT NULL REFERENCES conversations(id) ON DELETE CASCADE,
|
||||
role TEXT NOT NULL CHECK(role IN ('user', 'assistant')),
|
||||
content TEXT NOT NULL,
|
||||
selected_version INTEGER NOT NULL DEFAULT 0,
|
||||
created_at TEXT NOT NULL
|
||||
);
|
||||
|
||||
CREATE INDEX IF NOT EXISTS idx_messages_conversation
|
||||
ON messages(conversation_id, id);
|
||||
|
||||
CREATE TABLE IF NOT EXISTS message_versions (
|
||||
message_id INTEGER NOT NULL REFERENCES messages(id) ON DELETE CASCADE,
|
||||
version_index INTEGER NOT NULL CHECK(version_index >= 0),
|
||||
content TEXT NOT NULL,
|
||||
PRIMARY KEY(message_id, version_index)
|
||||
);
|
||||
"""
|
||||
)
|
||||
message_columns = {
|
||||
row["name"] for row in connection.execute("PRAGMA table_info(messages)")
|
||||
}
|
||||
if "selected_version" not in message_columns:
|
||||
connection.execute(
|
||||
"ALTER TABLE messages ADD COLUMN selected_version INTEGER NOT NULL DEFAULT 0"
|
||||
)
|
||||
|
||||
|
||||
def ensure_user(user_id: str) -> None:
|
||||
@@ -113,18 +134,47 @@ def get_conversation(user_id: str, conversation_id: str) -> dict[str, Any] | Non
|
||||
if conversation is None:
|
||||
return None
|
||||
messages = connection.execute(
|
||||
"SELECT role, content, created_at FROM messages "
|
||||
"SELECT id, role, content, selected_version, created_at FROM messages "
|
||||
"WHERE conversation_id = ? ORDER BY id",
|
||||
(conversation_id,),
|
||||
).fetchall()
|
||||
return {**dict(conversation), "messages": [dict(row) for row in messages]}
|
||||
versions = connection.execute(
|
||||
"""
|
||||
SELECT mv.message_id, mv.content
|
||||
FROM message_versions AS mv
|
||||
JOIN messages AS m ON m.id = mv.message_id
|
||||
WHERE m.conversation_id = ?
|
||||
ORDER BY mv.message_id, mv.version_index
|
||||
""",
|
||||
(conversation_id,),
|
||||
).fetchall()
|
||||
versions_by_message: dict[int, list[str]] = {}
|
||||
for version in versions:
|
||||
versions_by_message.setdefault(version["message_id"], []).append(
|
||||
version["content"]
|
||||
)
|
||||
saved_messages = []
|
||||
for message in messages:
|
||||
message_versions = versions_by_message.get(message["id"], [])
|
||||
if message["role"] == "assistant" and not message_versions:
|
||||
message_versions = [message["content"]]
|
||||
saved_messages.append(
|
||||
{
|
||||
"role": message["role"],
|
||||
"content": message["content"],
|
||||
"selected_version": message["selected_version"],
|
||||
"versions": message_versions,
|
||||
"created_at": message["created_at"],
|
||||
}
|
||||
)
|
||||
return {**dict(conversation), "messages": saved_messages}
|
||||
|
||||
|
||||
def save_conversation(
|
||||
user_id: str,
|
||||
conversation_id: str,
|
||||
title: str,
|
||||
messages: list[dict[str, str]],
|
||||
messages: list[dict[str, Any]],
|
||||
timestamp: str,
|
||||
) -> None:
|
||||
with _connect() as connection:
|
||||
@@ -145,13 +195,33 @@ def save_conversation(
|
||||
(conversation_id, user_id, title, timestamp, timestamp),
|
||||
)
|
||||
connection.execute("DELETE FROM messages WHERE conversation_id = ?", (conversation_id,))
|
||||
connection.executemany(
|
||||
"INSERT INTO messages(conversation_id, role, content, created_at) VALUES (?, ?, ?, ?)",
|
||||
[
|
||||
(conversation_id, message["role"], message["content"], timestamp)
|
||||
for message in messages
|
||||
],
|
||||
)
|
||||
for message in messages:
|
||||
cursor = connection.execute(
|
||||
"""
|
||||
INSERT INTO messages(
|
||||
conversation_id, role, content, selected_version, created_at
|
||||
) VALUES (?, ?, ?, ?, ?)
|
||||
""",
|
||||
(
|
||||
conversation_id,
|
||||
message["role"],
|
||||
message["content"],
|
||||
message.get("selected_version", 0),
|
||||
timestamp,
|
||||
),
|
||||
)
|
||||
if message["role"] == "assistant":
|
||||
versions = message.get("versions") or [message["content"]]
|
||||
connection.executemany(
|
||||
"""
|
||||
INSERT INTO message_versions(message_id, version_index, content)
|
||||
VALUES (?, ?, ?)
|
||||
""",
|
||||
[
|
||||
(cursor.lastrowid, index, content)
|
||||
for index, content in enumerate(versions)
|
||||
],
|
||||
)
|
||||
|
||||
|
||||
def delete_conversation(user_id: str, conversation_id: str) -> bool:
|
||||
|
||||
@@ -14,7 +14,7 @@ import httpx
|
||||
from fastapi import FastAPI, File, Form, Header, HTTPException, UploadFile
|
||||
from fastapi.middleware.cors import CORSMiddleware
|
||||
from fastapi.responses import StreamingResponse
|
||||
from pydantic import BaseModel, Field
|
||||
from pydantic import BaseModel, Field, model_validator
|
||||
|
||||
from app import database
|
||||
from app import workspace
|
||||
@@ -69,6 +69,21 @@ class WorkspaceAgentRequest(AgentRequest):
|
||||
class StoredMessage(BaseModel):
|
||||
role: str = Field(pattern="^(user|assistant)$")
|
||||
content: str
|
||||
versions: list[str] = Field(default_factory=list, max_length=32)
|
||||
selected_version: int = Field(default=0, ge=0)
|
||||
|
||||
@model_validator(mode="after")
|
||||
def validate_answer_versions(self) -> "StoredMessage":
|
||||
if self.role == "user" and self.versions:
|
||||
raise ValueError("User messages cannot contain assistant answer versions.")
|
||||
if self.versions:
|
||||
if self.selected_version >= len(self.versions):
|
||||
raise ValueError("selected_version is outside the versions list.")
|
||||
if self.versions[self.selected_version] != self.content:
|
||||
raise ValueError("content must match the selected answer version.")
|
||||
elif self.selected_version != 0:
|
||||
raise ValueError("selected_version must be zero when versions are omitted.")
|
||||
return self
|
||||
|
||||
|
||||
class ConversationWrite(BaseModel):
|
||||
|
||||
Reference in New Issue
Block a user