Add file analysis and answer feedback

This commit is contained in:
Hamza Ayed
2026-10-01 15:13:07 +03:00
parent dc4d5bad0f
commit 3563a104a3
10 changed files with 586 additions and 89 deletions
+88 -2
View File
@@ -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()