""" ByteAstra — Session management service using PyMongo. """ from __future__ import annotations import json import logging import uuid from datetime import datetime, timezone from pymongo.database import Database from app.models import Session, Message, MessageRole from app.schemas import Citation logger = logging.getLogger(__name__) # ── Session operations ───────────────────────────────────────────────────────── def create_session(db: Database, domain: str, title: str = "New Session", user_id: str | None = None) -> Session: doc = { "_id": str(uuid.uuid4()), "domain": domain, "title": title, "user_id": user_id, "subscription_tier": "free", "created_at": datetime.now(timezone.utc), "updated_at": datetime.now(timezone.utc), } db.sessions.insert_one(doc) logger.info("Created session %s (domain=%s) in MongoDB", doc["_id"], domain) return Session(doc, []) def get_session(db: Database, session_id: str) -> Session | None: sess_doc = db.sessions.find_one({"_id": session_id}) if not sess_doc: return None # Retrieve messages sorted by created_at ascending msg_docs = db.messages.find({"session_id": session_id}).sort("created_at", 1) messages = [Message(m) for m in msg_docs] return Session(sess_doc, messages) def list_sessions(db: Database, user_id: str | None = None, limit: int = 50) -> list[Session]: query = {} if user_id: query["user_id"] = user_id sess_docs = db.sessions.find(query).sort("updated_at", -1).limit(limit) return [Session(s, []) for s in sess_docs] def delete_session(db: Database, session_id: str) -> bool: res = db.sessions.delete_one({"_id": session_id}) if res.deleted_count == 0: return False # Cascading delete of messages db.messages.delete_many({"session_id": session_id}) return True def update_session_title(db: Database, session_id: str, title: str) -> Session | None: from pymongo import ReturnDocument res = db.sessions.find_one_and_update( {"_id": session_id}, {"$set": {"title": title, "updated_at": datetime.now(timezone.utc)}}, return_document=ReturnDocument.AFTER ) if not res: return None msg_docs = db.messages.find({"session_id": session_id}).sort("created_at", 1) messages = [Message(m) for m in msg_docs] return Session(res, messages) # ── Message operations ───────────────────────────────────────────────────────── def add_user_message(db: Database, session_id: str, content: str) -> Message: doc = { "_id": str(uuid.uuid4()), "session_id": session_id, "role": MessageRole.user.value, "content": content, "created_at": datetime.now(timezone.utc), } db.messages.insert_one(doc) db.sessions.update_one( {"_id": session_id}, {"$set": {"updated_at": datetime.now(timezone.utc)}} ) return Message(doc) def add_assistant_message( db: Database, session_id: str, content: str, citations: list[Citation], is_grounded: bool, ) -> Message: doc = { "_id": str(uuid.uuid4()), "session_id": session_id, "role": MessageRole.assistant.value, "content": content, "citations_json": json.dumps([c.model_dump() for c in citations]), "is_grounded": is_grounded, "created_at": datetime.now(timezone.utc), } db.messages.insert_one(doc) db.sessions.update_one( {"_id": session_id}, {"$set": {"updated_at": datetime.now(timezone.utc)}} ) return Message(doc) def get_recent_history( db: Database, session_id: str, max_turns: int = 6 ) -> list[dict]: """ Return recent messages as a list of OpenAI-format dicts for the prompt. """ msg_docs = db.messages.find( {"session_id": session_id, "role": {"$ne": "system"}} ).sort("created_at", -1).limit(max_turns * 2) # Reverse to chronological order msg_list = list(reversed(list(msg_docs))) return [{"role": m["role"], "content": m["content"]} for m in msg_list] def parse_citations(message: Message) -> list[Citation]: if not message.citations_json: return [] try: return [Citation(**c) for c in json.loads(message.citations_json)] except Exception: return []