Mendu9 commited on
Commit ·
eafc193
1
Parent(s): 7b66c45
feat: add SQLAlchemy models for all 5 DB tables
Browse filesTDD implementation of ChatSession, ResponseMetrics, ResponseFeedback,
LLMJudgeScore, and GuardrailEvent models. Uses generic JSON/Uuid/Integer
types for SQLite test compatibility while remaining PostgreSQL-ready.
- mao/db/__init__.py +35 -0
- mao/db/models.py +80 -0
- tests/db/__init__.py +0 -0
- tests/db/test_models.py +58 -0
mao/db/__init__.py
ADDED
|
@@ -0,0 +1,35 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
from __future__ import annotations
|
| 2 |
+
import logging
|
| 3 |
+
from contextlib import contextmanager
|
| 4 |
+
from sqlalchemy import create_engine
|
| 5 |
+
from sqlalchemy.orm import Session, sessionmaker
|
| 6 |
+
from mao.db.models import Base
|
| 7 |
+
|
| 8 |
+
logger = logging.getLogger(__name__)
|
| 9 |
+
|
| 10 |
+
_engine = None
|
| 11 |
+
_SessionLocal = None
|
| 12 |
+
|
| 13 |
+
POSTGRES_URL = "postgresql://mao:mao@localhost:5432/mao"
|
| 14 |
+
|
| 15 |
+
|
| 16 |
+
def init_db() -> None:
|
| 17 |
+
global _engine, _SessionLocal
|
| 18 |
+
try:
|
| 19 |
+
_engine = create_engine(POSTGRES_URL, pool_pre_ping=True)
|
| 20 |
+
Base.metadata.create_all(_engine)
|
| 21 |
+
_SessionLocal = sessionmaker(bind=_engine)
|
| 22 |
+
logger.info("Database initialised: %s", POSTGRES_URL)
|
| 23 |
+
except Exception as exc:
|
| 24 |
+
logger.critical("Database init failed: %s", exc)
|
| 25 |
+
|
| 26 |
+
|
| 27 |
+
@contextmanager
|
| 28 |
+
def get_db_session():
|
| 29 |
+
if _SessionLocal is None:
|
| 30 |
+
raise RuntimeError("Database not initialised — call init_db() first")
|
| 31 |
+
session = _SessionLocal()
|
| 32 |
+
try:
|
| 33 |
+
yield session
|
| 34 |
+
finally:
|
| 35 |
+
session.close()
|
mao/db/models.py
ADDED
|
@@ -0,0 +1,80 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
from __future__ import annotations
|
| 2 |
+
import uuid
|
| 3 |
+
from datetime import datetime, timezone
|
| 4 |
+
from sqlalchemy import BigInteger, Boolean, Column, Float, ForeignKey, Integer, JSON, SmallInteger, String, Text, Uuid
|
| 5 |
+
from sqlalchemy.orm import DeclarativeBase
|
| 6 |
+
|
| 7 |
+
|
| 8 |
+
class Base(DeclarativeBase):
|
| 9 |
+
pass
|
| 10 |
+
|
| 11 |
+
|
| 12 |
+
def _uuid():
|
| 13 |
+
return uuid.uuid4()
|
| 14 |
+
|
| 15 |
+
|
| 16 |
+
def _now():
|
| 17 |
+
return datetime.now(timezone.utc)
|
| 18 |
+
|
| 19 |
+
|
| 20 |
+
class ChatSession(Base):
|
| 21 |
+
__tablename__ = "chat_sessions"
|
| 22 |
+
session_id = Column(Uuid(as_uuid=True), primary_key=True, default=_uuid)
|
| 23 |
+
user_id = Column(String, nullable=False)
|
| 24 |
+
user_query = Column(Text, nullable=False)
|
| 25 |
+
pii_scrubbed_query = Column(Text)
|
| 26 |
+
response = Column(Text)
|
| 27 |
+
agent_used = Column(String)
|
| 28 |
+
domain = Column(String)
|
| 29 |
+
report_card = Column(JSON)
|
| 30 |
+
council_verdict = Column(JSON)
|
| 31 |
+
nli_flags = Column(JSON)
|
| 32 |
+
uncertainty_flag = Column(Boolean, default=False)
|
| 33 |
+
created_at = Column(Text, default=lambda: _now().isoformat())
|
| 34 |
+
|
| 35 |
+
|
| 36 |
+
class ResponseMetrics(Base):
|
| 37 |
+
__tablename__ = "response_metrics"
|
| 38 |
+
id = Column(Integer, primary_key=True, autoincrement=True)
|
| 39 |
+
request_id = Column(String, nullable=False)
|
| 40 |
+
session_id = Column(Uuid(as_uuid=True), ForeignKey("chat_sessions.session_id"))
|
| 41 |
+
user_id = Column(String, nullable=False)
|
| 42 |
+
agent_used = Column(String)
|
| 43 |
+
faithfulness = Column(Float)
|
| 44 |
+
answer_relevancy = Column(Float)
|
| 45 |
+
context_precision = Column(Float)
|
| 46 |
+
context_recall = Column(Float)
|
| 47 |
+
latency_ms = Column(Float)
|
| 48 |
+
created_at = Column(Text, default=lambda: _now().isoformat())
|
| 49 |
+
|
| 50 |
+
|
| 51 |
+
class ResponseFeedback(Base):
|
| 52 |
+
__tablename__ = "response_feedback"
|
| 53 |
+
id = Column(Integer, primary_key=True, autoincrement=True)
|
| 54 |
+
session_id = Column(Uuid(as_uuid=True), ForeignKey("chat_sessions.session_id"))
|
| 55 |
+
user_id = Column(String, nullable=False)
|
| 56 |
+
rating = Column(SmallInteger, nullable=False)
|
| 57 |
+
comment = Column(Text)
|
| 58 |
+
created_at = Column(Text, default=lambda: _now().isoformat())
|
| 59 |
+
|
| 60 |
+
|
| 61 |
+
class LLMJudgeScore(Base):
|
| 62 |
+
__tablename__ = "llm_judge_scores"
|
| 63 |
+
id = Column(Integer, primary_key=True, autoincrement=True)
|
| 64 |
+
session_id = Column(Uuid(as_uuid=True), ForeignKey("chat_sessions.session_id"))
|
| 65 |
+
accuracy = Column(SmallInteger)
|
| 66 |
+
completeness = Column(SmallInteger)
|
| 67 |
+
safety = Column(SmallInteger)
|
| 68 |
+
clarity = Column(SmallInteger)
|
| 69 |
+
notes = Column(Text)
|
| 70 |
+
created_at = Column(Text, default=lambda: _now().isoformat())
|
| 71 |
+
|
| 72 |
+
|
| 73 |
+
class GuardrailEvent(Base):
|
| 74 |
+
__tablename__ = "guardrail_events"
|
| 75 |
+
id = Column(Integer, primary_key=True, autoincrement=True)
|
| 76 |
+
session_id = Column(Uuid(as_uuid=True))
|
| 77 |
+
guardrail_name = Column(String, nullable=False)
|
| 78 |
+
triggered = Column(Boolean, nullable=False)
|
| 79 |
+
detail = Column(Text)
|
| 80 |
+
created_at = Column(Text, default=lambda: _now().isoformat())
|
tests/db/__init__.py
ADDED
|
File without changes
|
tests/db/test_models.py
ADDED
|
@@ -0,0 +1,58 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import pytest
|
| 2 |
+
import uuid
|
| 3 |
+
from sqlalchemy import create_engine
|
| 4 |
+
from sqlalchemy.orm import Session
|
| 5 |
+
from mao.db.models import (
|
| 6 |
+
Base, ChatSession, ResponseMetrics,
|
| 7 |
+
ResponseFeedback, LLMJudgeScore, GuardrailEvent,
|
| 8 |
+
)
|
| 9 |
+
|
| 10 |
+
@pytest.fixture
|
| 11 |
+
def engine():
|
| 12 |
+
eng = create_engine("sqlite:///:memory:")
|
| 13 |
+
Base.metadata.create_all(eng)
|
| 14 |
+
yield eng
|
| 15 |
+
Base.metadata.drop_all(eng)
|
| 16 |
+
|
| 17 |
+
@pytest.fixture
|
| 18 |
+
def session(engine):
|
| 19 |
+
with Session(engine) as s:
|
| 20 |
+
yield s
|
| 21 |
+
|
| 22 |
+
def test_chat_session_insert(session):
|
| 23 |
+
row = ChatSession(session_id=uuid.uuid4(), user_id="user-1", user_query="What is amyloid?")
|
| 24 |
+
session.add(row)
|
| 25 |
+
session.commit()
|
| 26 |
+
result = session.get(ChatSession, row.session_id)
|
| 27 |
+
assert result.user_id == "user-1"
|
| 28 |
+
assert result.uncertainty_flag is False
|
| 29 |
+
|
| 30 |
+
def test_response_metrics_insert(session):
|
| 31 |
+
sid = uuid.uuid4()
|
| 32 |
+
session.add(ChatSession(session_id=sid, user_id="u", user_query="q"))
|
| 33 |
+
session.commit()
|
| 34 |
+
row = ResponseMetrics(request_id="req-1", session_id=sid, user_id="u", faithfulness=0.85)
|
| 35 |
+
session.add(row)
|
| 36 |
+
session.commit()
|
| 37 |
+
assert session.get(ResponseMetrics, row.id).faithfulness == pytest.approx(0.85)
|
| 38 |
+
|
| 39 |
+
def test_response_feedback_insert(session):
|
| 40 |
+
sid = uuid.uuid4()
|
| 41 |
+
session.add(ChatSession(session_id=sid, user_id="u", user_query="q"))
|
| 42 |
+
session.commit()
|
| 43 |
+
session.add(ResponseFeedback(session_id=sid, user_id="u", rating=1))
|
| 44 |
+
session.commit()
|
| 45 |
+
assert session.query(ResponseFeedback).first().rating == 1
|
| 46 |
+
|
| 47 |
+
def test_llm_judge_scores_insert(session):
|
| 48 |
+
sid = uuid.uuid4()
|
| 49 |
+
session.add(ChatSession(session_id=sid, user_id="u", user_query="q"))
|
| 50 |
+
session.commit()
|
| 51 |
+
session.add(LLMJudgeScore(session_id=sid, accuracy=8, safety=10))
|
| 52 |
+
session.commit()
|
| 53 |
+
assert session.query(LLMJudgeScore).first().accuracy == 8
|
| 54 |
+
|
| 55 |
+
def test_guardrail_event_insert(session):
|
| 56 |
+
session.add(GuardrailEvent(guardrail_name="token_limit", triggered=True, detail="550 tokens"))
|
| 57 |
+
session.commit()
|
| 58 |
+
assert session.query(GuardrailEvent).first().triggered is True
|