byteastra / app /services /sessions.py
risu1012's picture
feat: optimize deployment payload using zipped database index
9fca47f
Raw
History Blame Contribute Delete
4.62 kB
"""
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 []