Add file analysis and answer feedback
This commit is contained in:
@@ -3,6 +3,7 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
import hashlib
|
||||
import sqlite3
|
||||
from uuid import UUID
|
||||
from pathlib import Path
|
||||
@@ -89,6 +90,19 @@ def initialize_database() -> None:
|
||||
content TEXT NOT NULL,
|
||||
PRIMARY KEY(message_id, version_index)
|
||||
);
|
||||
|
||||
CREATE TABLE IF NOT EXISTS answer_feedback (
|
||||
user_id TEXT NOT NULL REFERENCES users(id) ON DELETE CASCADE,
|
||||
conversation_id TEXT NOT NULL REFERENCES conversations(id) ON DELETE CASCADE,
|
||||
message_index INTEGER NOT NULL CHECK(message_index >= 0),
|
||||
version_index INTEGER NOT NULL CHECK(version_index >= 0),
|
||||
answer_hash TEXT NOT NULL,
|
||||
rating INTEGER NOT NULL CHECK(rating IN (-1, 1)),
|
||||
updated_at TEXT NOT NULL DEFAULT CURRENT_TIMESTAMP,
|
||||
PRIMARY KEY(user_id, conversation_id, message_index, version_index, answer_hash)
|
||||
);
|
||||
CREATE INDEX IF NOT EXISTS idx_answer_feedback_user
|
||||
ON answer_feedback(user_id, updated_at DESC);
|
||||
"""
|
||||
)
|
||||
message_columns = {
|
||||
@@ -148,13 +162,25 @@ def get_conversation(user_id: str, conversation_id: str) -> dict[str, Any] | Non
|
||||
""",
|
||||
(conversation_id,),
|
||||
).fetchall()
|
||||
versions_by_message: dict[int, list[str]] = {}
|
||||
versions_by_message: dict[int, list[str]] = {}
|
||||
for version in versions:
|
||||
versions_by_message.setdefault(version["message_id"], []).append(
|
||||
version["content"]
|
||||
)
|
||||
feedback_by_answer: dict[tuple[int, int, str], int] = {}
|
||||
with _connect() as connection:
|
||||
feedback_rows = connection.execute(
|
||||
"""
|
||||
SELECT message_index, version_index, answer_hash, rating
|
||||
FROM answer_feedback
|
||||
WHERE user_id = ? AND conversation_id = ?
|
||||
""",
|
||||
(user_id, conversation_id),
|
||||
).fetchall()
|
||||
for row in feedback_rows:
|
||||
feedback_by_answer[(row["message_index"], row["version_index"], row["answer_hash"])] = row["rating"]
|
||||
saved_messages = []
|
||||
for message in messages:
|
||||
for message_index, message in enumerate(messages):
|
||||
message_versions = versions_by_message.get(message["id"], [])
|
||||
if message["role"] == "assistant" and not message_versions:
|
||||
message_versions = [message["content"]]
|
||||
@@ -164,6 +190,20 @@ def get_conversation(user_id: str, conversation_id: str) -> dict[str, Any] | Non
|
||||
"content": message["content"],
|
||||
"selected_version": message["selected_version"],
|
||||
"versions": message_versions,
|
||||
"feedback_versions": (
|
||||
[
|
||||
feedback_by_answer.get(
|
||||
(
|
||||
message_index,
|
||||
version_index,
|
||||
hashlib.sha256(content.encode("utf-8")).hexdigest(),
|
||||
)
|
||||
)
|
||||
for version_index, content in enumerate(message_versions)
|
||||
]
|
||||
if message["role"] == "assistant"
|
||||
else []
|
||||
),
|
||||
"created_at": message["created_at"],
|
||||
}
|
||||
)
|
||||
@@ -233,4 +273,50 @@ def delete_conversation(user_id: str, conversation_id: str) -> bool:
|
||||
return cursor.rowcount > 0
|
||||
|
||||
|
||||
def save_answer_feedback(
|
||||
user_id: str,
|
||||
conversation_id: str,
|
||||
message_index: int,
|
||||
version_index: int,
|
||||
rating: int,
|
||||
timestamp: str,
|
||||
) -> None:
|
||||
"""Store a user's rating for the exact assistant answer version."""
|
||||
with _connect() as connection:
|
||||
messages = connection.execute(
|
||||
"SELECT id, role, content FROM messages WHERE conversation_id = ? ORDER BY id",
|
||||
(conversation_id,),
|
||||
).fetchall()
|
||||
if message_index >= len(messages):
|
||||
raise ValueError("Message index does not exist.")
|
||||
message = messages[message_index]
|
||||
if message["role"] != "assistant":
|
||||
raise ValueError("Feedback can only be attached to assistant messages.")
|
||||
version = connection.execute(
|
||||
"SELECT content FROM message_versions WHERE message_id = ? AND version_index = ?",
|
||||
(message["id"], version_index),
|
||||
).fetchone()
|
||||
if version is None:
|
||||
if version_index != 0:
|
||||
raise ValueError("Answer version does not exist.")
|
||||
answer = message["content"]
|
||||
else:
|
||||
answer = version["content"]
|
||||
answer_hash = hashlib.sha256(answer.encode("utf-8")).hexdigest()
|
||||
connection.execute(
|
||||
"""
|
||||
INSERT INTO answer_feedback(
|
||||
user_id, conversation_id, message_index, version_index,
|
||||
answer_hash, rating, updated_at
|
||||
) VALUES (?, ?, ?, ?, ?, ?, ?)
|
||||
ON CONFLICT(user_id, conversation_id, message_index, version_index, answer_hash)
|
||||
DO UPDATE SET rating = excluded.rating, updated_at = excluded.updated_at
|
||||
""",
|
||||
(
|
||||
user_id, conversation_id, message_index, version_index,
|
||||
answer_hash, rating, timestamp,
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
initialize_database()
|
||||
|
||||
Reference in New Issue
Block a user