Mendu9 commited on
Commit
eafc193
·
1 Parent(s): 7b66c45

feat: add SQLAlchemy models for all 5 DB tables

Browse files

TDD 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 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