from __future__ import annotations from contextlib import contextmanager from datetime import UTC, datetime from typing import Any, Iterator from uuid import uuid4 from sqlalchemy import Connection, Engine, create_engine, delete, func, insert, select, update from .db.schema import build_session_tables from .models import BusinessResume, ComponentBlock, ConversationTurn, SessionView from .resume_document_core import attach_gap_report_staleness def _now() -> datetime: return datetime.now(UTC) class PostgresDatabase: """PostgreSQL implementation of the Resume Agent persistence contract.""" def __init__(self, database_url: str, *, schema: str = "resume_agent") -> None: self.engine: Engine = create_engine(database_url, pool_pre_ping=True) self.schema = schema self.metadata, self.tables = build_session_tables(schema) @contextmanager def transaction(self, *, immediate: bool = False) -> Iterator[Connection]: del immediate with self.engine.begin() as connection: yield connection def initialize(self) -> None: with self.engine.begin() as connection: connection.exec_driver_sql(f'CREATE SCHEMA IF NOT EXISTS "{self.schema}"') self.metadata.create_all(connection) def create_session( self, session_id: str, stage: str, profile: dict[str, Any], initial_turn: dict[str, Any] ) -> None: now = _now() with self.transaction() as connection: connection.execute(insert(self.tables["sessions"]).values( id=session_id, stage=stage, revision=0, profile=profile, created_at=now, updated_at=now, )) self.insert_turn(connection, session_id=session_id, **initial_turn) def fetch_session( self, connection: Connection, session_id: str, *, for_update: bool = False, ) -> dict[str, Any] | None: sessions = self.tables["sessions"] statement = select(sessions).where(sessions.c.id == session_id) if for_update: statement = statement.with_for_update() row = connection.execute(statement).mappings().first() if row is None: return None result = dict(row) profile = dict(result["profile"]) if profile.get("job_type") == "other": profile["job_type"] = "internship" result["profile"] = profile result["updated_at"] = _now() connection.execute(update(sessions).where(sessions.c.id == session_id).values( profile=profile, updated_at=result["updated_at"] )) return result def get_session(self, session_id: str) -> dict[str, Any] | None: with self.transaction() as connection: return self.fetch_session(connection, session_id) def find_latest_session_by_external_user_id( self, external_user_id: str ) -> dict[str, Any] | None: normalized_user_id = str(external_user_id or "").strip() if not normalized_user_id: return None sessions = self.tables["sessions"] with self.transaction() as connection: session_id = connection.execute( select(sessions.c.id) .where( func.jsonb_extract_path_text( sessions.c.profile, "external_account", "provider" ) == "offerpai", func.jsonb_extract_path_text( sessions.c.profile, "external_account", "user_id" ) == normalized_user_id, ) .order_by( sessions.c.updated_at.desc(), sessions.c.created_at.desc(), sessions.c.id.desc(), ) .limit(1) ).scalar_one_or_none() return ( self.fetch_session(connection, session_id) if session_id is not None else None ) def update_session( self, connection: Connection, session_id: str, *, stage: str, profile: dict[str, Any], draft_id: str | None = None, resume_id: str | None = None, increment_revision: bool = True, ) -> dict[str, Any]: sessions = self.tables["sessions"] current = connection.execute( select(sessions).where(sessions.c.id == session_id).with_for_update() ).mappings().first() if current is None: raise KeyError(session_id) values = { "stage": stage, "profile": profile, "revision": current["revision"] + (1 if increment_revision else 0), "draft_id": draft_id if draft_id is not None else current["draft_id"], "resume_id": resume_id if resume_id is not None else current["resume_id"], "updated_at": _now(), } connection.execute(update(sessions).where(sessions.c.id == session_id).values(**values)) return self.fetch_session(connection, session_id) # type: ignore[return-value] def insert_turn( self, connection: Connection, *, session_id: str, role: str, content: str | None, composer_mode: str, blocks: list[dict[str, Any]], ) -> str: sessions, turns, block_table = ( self.tables["sessions"], self.tables["turns"], self.tables["blocks"] ) connection.execute(select(sessions.c.id).where(sessions.c.id == session_id).with_for_update()).one() sequence = connection.execute( select(func.coalesce(func.max(turns.c.sequence), 0) + 1).where(turns.c.session_id == session_id) ).scalar_one() now, turn_id = _now(), f"turn_{uuid4().hex}" connection.execute(insert(turns).values( id=turn_id, session_id=session_id, sequence=sequence, role=role, content=content, composer_mode=composer_mode, created_at=now, )) if blocks: connection.execute(insert(block_table), [{ "id": block.get("id", f"block_{uuid4().hex}"), "session_id": session_id, "turn_id": turn_id, "block_index": index, "type": block["type"], "lifecycle": block.get("lifecycle", "active"), "data": block.get("data", {}), "version": 1, "created_at": now, "updated_at": now, } for index, block in enumerate(blocks)]) return turn_id def fetch_block( self, connection: Connection, session_id: str, block_id: str ) -> dict[str, Any] | None: blocks = self.tables["blocks"] row = connection.execute(select(blocks).where( blocks.c.id == block_id, blocks.c.session_id == session_id )).mappings().first() return dict(row) if row else None def update_block( self, connection: Connection, block_id: str, *, lifecycle: str, data: dict[str, Any] | None = None, ) -> None: blocks = self.tables["blocks"] current = connection.execute( select(blocks.c.data).where(blocks.c.id == block_id).with_for_update() ).first() if current is None: raise KeyError(block_id) connection.execute(update(blocks).where(blocks.c.id == block_id).values( lifecycle=lifecycle, data=current._mapping["data"] if data is None else data, version=blocks.c.version + 1, updated_at=_now(), )) def supersede_active_components(self, connection: Connection, session_id: str) -> None: blocks = self.tables["blocks"] connection.execute(update(blocks).where( blocks.c.session_id == session_id, blocks.c.type == "component", blocks.c.lifecycle == "active", ).values(lifecycle="superseded", version=blocks.c.version + 1, updated_at=_now())) def get_turn(self, turn_id: str) -> ConversationTurn: with self.transaction() as connection: return self.fetch_turn(connection, turn_id) def fetch_turn(self, connection: Connection, turn_id: str) -> ConversationTurn: turns = self.tables["turns"] row = connection.execute(select(turns).where(turns.c.id == turn_id)).mappings().first() if row is None: raise KeyError(turn_id) return self._turn_from_row(connection, row) def list_turns(self, session_id: str) -> list[ConversationTurn]: turns = self.tables["turns"] with self.transaction() as connection: rows = connection.execute(select(turns).where( turns.c.session_id == session_id ).order_by(turns.c.sequence)).mappings().all() return [self._turn_from_row(connection, row) for row in rows] def _turn_from_row(self, connection: Connection, row: Any) -> ConversationTurn: blocks = self.tables["blocks"] block_rows = connection.execute(select(blocks).where( blocks.c.turn_id == row["id"] ).order_by(blocks.c.block_index)).mappings().all() return ConversationTurn( id=row["id"], sequence=row["sequence"], role=row["role"], content=row["content"], composer_mode=row["composer_mode"], blocks=[ComponentBlock( id=block["id"], type=block["type"], lifecycle=block["lifecycle"], data=block["data"], version=block["version"], created_at=block["created_at"], updated_at=block["updated_at"], ) for block in block_rows], created_at=row["created_at"], ) def session_view(self, session: dict[str, Any]) -> SessionView: profile = session["profile"] phone = profile.get("phone") return SessionView( id=session["id"], stage=session["stage"], revision=session["revision"], job_type=profile.get("job_type"), anchor_type=profile.get("anchor_type"), masked_phone=f"{phone[:3]}****{phone[-4:]}" if phone else None, phone_source=profile.get("phone_source"), name=profile.get("name"), draft_id=session.get("draft_id"), resume_id=session.get("resume_id"), created_at=session["created_at"], updated_at=session["updated_at"], ) def fetch_resume(self, connection: Connection, session_id: str) -> dict[str, Any] | None: resumes = self.tables["resumes"] row = connection.execute(select(resumes).where( resumes.c.session_id == session_id )).mappings().first() return dict(row) if row else None def insert_resume( self, connection: Connection, *, resume_id: str, session_id: str, idempotency_key: str | None, content: dict[str, Any], ) -> dict[str, Any]: now = _now() connection.execute(insert(self.tables["resumes"]).values( id=resume_id, session_id=session_id, idempotency_key=idempotency_key, revision=1, content=content, created_at=now, updated_at=now, )) return self.fetch_resume(connection, session_id) # type: ignore[return-value] def update_resume( self, connection: Connection, session_id: str, content: dict[str, Any] ) -> dict[str, Any]: resumes = self.tables["resumes"] result = connection.execute(update(resumes).where( resumes.c.session_id == session_id ).values(content=content, revision=resumes.c.revision + 1, updated_at=_now())) if result.rowcount != 1: raise KeyError(session_id) return self.fetch_resume(connection, session_id) # type: ignore[return-value] def create_optimization_run( self, connection: Connection, *, run_id: str, session_id: str, entry_id: str, mode: str, status: str, source_revision: int, state: dict[str, Any], proposal: dict[str, Any] | None = None, ) -> dict[str, Any]: now = _now() connection.execute(insert(self.tables["optimization_runs"]).values( id=run_id, session_id=session_id, entry_id=entry_id, mode=mode, status=status, source_revision=source_revision, state=state, proposal=proposal, created_at=now, updated_at=now, )) return self.fetch_optimization_run(connection, session_id, run_id) # type: ignore[return-value] def fetch_optimization_run( self, connection: Connection, session_id: str, run_id: str ) -> dict[str, Any] | None: runs = self.tables["optimization_runs"] row = connection.execute(select(runs).where( runs.c.id == run_id, runs.c.session_id == session_id )).mappings().first() return dict(row) if row else None def find_active_optimization_run( self, connection: Connection, session_id: str, entry_id: str ) -> dict[str, Any] | None: runs = self.tables["optimization_runs"] row = connection.execute(select(runs).where( runs.c.session_id == session_id, runs.c.entry_id == entry_id, runs.c.status.in_(("question_pending", "proposal_pending")), ).order_by(runs.c.created_at.desc()).limit(1)).mappings().first() return dict(row) if row else None def list_active_optimization_runs( self, connection: Connection, session_id: str ) -> list[dict[str, Any]]: runs = self.tables["optimization_runs"] rows = connection.execute(select(runs).where( runs.c.session_id == session_id, runs.c.status.in_(("question_pending", "proposal_pending")), ).order_by(runs.c.updated_at, runs.c.created_at)).mappings().all() return [dict(row) for row in rows] def update_optimization_run( self, connection: Connection, *, session_id: str, run_id: str, status: str, state: dict[str, Any], proposal: dict[str, Any] | None, ) -> dict[str, Any]: runs = self.tables["optimization_runs"] result = connection.execute(update(runs).where( runs.c.id == run_id, runs.c.session_id == session_id ).values(status=status, state=state, proposal=proposal, updated_at=_now())) if result.rowcount != 1: raise KeyError(run_id) return self.fetch_optimization_run(connection, session_id, run_id) # type: ignore[return-value] def create_resume_import( self, connection: Connection, *, import_id: str, session_id: str, file_name: str, mime_type: str, size_bytes: int, sha256: str, object_key: str, document: dict[str, Any], field_reviews: list[dict[str, Any]], ) -> dict[str, Any]: now = _now() connection.execute(insert(self.tables["resume_imports"]).values( id=import_id, session_id=session_id, file_name=file_name, mime_type=mime_type, size_bytes=size_bytes, sha256=sha256, object_key=object_key, status="awaiting_review", document=document, field_reviews=field_reviews, created_at=now, updated_at=now, )) return self.fetch_resume_import(connection, session_id, import_id) # type: ignore[return-value] def fetch_resume_import( self, connection: Connection, session_id: str, import_id: str ) -> dict[str, Any] | None: imports = self.tables["resume_imports"] row = connection.execute(select(imports).where( imports.c.id == import_id, imports.c.session_id == session_id )).mappings().first() return dict(row) if row else None def find_resume_import_by_sha256( self, connection: Connection, session_id: str, sha256: str ) -> dict[str, Any] | None: imports = self.tables["resume_imports"] row = connection.execute(select(imports).where( imports.c.session_id == session_id, imports.c.sha256 == sha256 )).mappings().first() return dict(row) if row else None def update_resume_import_status( self, connection: Connection, session_id: str, import_id: str, status: str ) -> dict[str, Any]: imports = self.tables["resume_imports"] result = connection.execute(update(imports).where( imports.c.id == import_id, imports.c.session_id == session_id ).values(status=status, updated_at=_now())) if result.rowcount != 1: raise KeyError(import_id) return self.fetch_resume_import(connection, session_id, import_id) # type: ignore[return-value] @staticmethod def resume_view(resume: dict[str, Any]) -> BusinessResume: return BusinessResume( id=resume["id"], session_id=resume["session_id"], revision=resume["revision"], content=attach_gap_report_staleness(resume["content"]), created_at=resume["created_at"], updated_at=resume["updated_at"], ) def delete_session(self, session_id: str) -> bool: with self.transaction() as connection: result = connection.execute(delete(self.tables["sessions"]).where( self.tables["sessions"].c.id == session_id )) return result.rowcount == 1