diff --git a/migrations/versions/a8f2d4e6c9b1_add_dialectic_traces_table.py b/migrations/versions/a8f2d4e6c9b1_add_dialectic_traces_table.py new file mode 100644 index 00000000..82c21aac --- /dev/null +++ b/migrations/versions/a8f2d4e6c9b1_add_dialectic_traces_table.py @@ -0,0 +1,119 @@ +"""add dialectic_traces table + +Revision ID: a8f2d4e6c9b1 +Revises: 7c0d9a4e3b1f +Create Date: 2026-01-29 + +""" + +from collections.abc import Sequence + +import sqlalchemy as sa +from alembic import op +from sqlalchemy.dialects import postgresql + +from migrations.utils import get_schema + +# revision identifiers, used by Alembic. +revision: str = "a8f2d4e6c9b1" +down_revision: str | None = "7c0d9a4e3b1f" +branch_labels: str | Sequence[str] | None = None +depends_on: str | Sequence[str] | None = None + +schema = get_schema() + + +def upgrade() -> None: + """Create the dialectic_traces table for internal logging of dialectic interactions.""" + op.create_table( + "dialectic_traces", + sa.Column("id", sa.TEXT(), nullable=False), + sa.Column("workspace_name", sa.TEXT(), nullable=False), + sa.Column("session_name", sa.TEXT(), nullable=True), + sa.Column("observer", sa.TEXT(), nullable=False), + sa.Column("observed", sa.TEXT(), nullable=False), + sa.Column("query", sa.TEXT(), nullable=False), + sa.Column( + "retrieved_doc_ids", + postgresql.JSONB(astext_type=sa.Text()), + server_default=sa.text("'[]'::jsonb"), + nullable=False, + ), + sa.Column( + "tool_calls", + postgresql.JSONB(astext_type=sa.Text()), + server_default=sa.text("'[]'::jsonb"), + nullable=False, + ), + sa.Column("response", sa.TEXT(), nullable=False), + sa.Column("reasoning_level", sa.TEXT(), nullable=False), + sa.Column("total_duration_ms", sa.Float(), nullable=False), + sa.Column("input_tokens", sa.Integer(), nullable=False), + sa.Column("output_tokens", sa.Integer(), nullable=False), + sa.Column( + "created_at", + sa.DateTime(timezone=True), + server_default=sa.text("now()"), + nullable=False, + ), + sa.ForeignKeyConstraint( + ["workspace_name"], + [f"{schema}.workspaces.name" if schema else "workspaces.name"], + name="fk_dialectic_traces_workspace_name", + ), + sa.PrimaryKeyConstraint("id"), + schema=schema, + ) + # Create indexes for efficient querying + op.create_index( + "ix_dialectic_traces_workspace_name", + "dialectic_traces", + ["workspace_name"], + schema=schema, + ) + op.create_index( + "ix_dialectic_traces_session_name", + "dialectic_traces", + ["session_name"], + schema=schema, + ) + op.create_index( + "ix_dialectic_traces_observer", + "dialectic_traces", + ["observer"], + schema=schema, + ) + op.create_index( + "ix_dialectic_traces_observed", + "dialectic_traces", + ["observed"], + schema=schema, + ) + op.create_index( + "ix_dialectic_traces_created_at", + "dialectic_traces", + ["created_at"], + schema=schema, + ) + + +def downgrade() -> None: + """Drop the dialectic_traces table.""" + op.drop_index( + "ix_dialectic_traces_created_at", table_name="dialectic_traces", schema=schema + ) + op.drop_index( + "ix_dialectic_traces_observed", table_name="dialectic_traces", schema=schema + ) + op.drop_index( + "ix_dialectic_traces_observer", table_name="dialectic_traces", schema=schema + ) + op.drop_index( + "ix_dialectic_traces_session_name", table_name="dialectic_traces", schema=schema + ) + op.drop_index( + "ix_dialectic_traces_workspace_name", + table_name="dialectic_traces", + schema=schema, + ) + op.drop_table("dialectic_traces", schema=schema) diff --git a/src/crud/__init__.py b/src/crud/__init__.py index ef54f599..3c126d43 100644 --- a/src/crud/__init__.py +++ b/src/crud/__init__.py @@ -1,5 +1,10 @@ from .collection import get_collection, get_or_create_collection from .deriver import get_deriver_status, get_queue_status +from .dialectic_trace import ( + create_dialectic_trace, + get_dialectic_trace_stats, + get_dialectic_traces, +) from .document import ( create_documents, create_observations, @@ -73,6 +78,10 @@ __all__ = [ # Deriver "get_deriver_status", "get_queue_status", + # Dialectic Trace + "create_dialectic_trace", + "get_dialectic_traces", + "get_dialectic_trace_stats", # Document "create_documents", "create_observations", diff --git a/src/crud/dialectic_trace.py b/src/crud/dialectic_trace.py new file mode 100644 index 00000000..e203cbbf --- /dev/null +++ b/src/crud/dialectic_trace.py @@ -0,0 +1,143 @@ +"""CRUD operations for DialecticTrace records.""" + +import datetime +import re +from typing import Any + +from sqlalchemy import func, select +from sqlalchemy.ext.asyncio import AsyncSession + +from src import models +from src.schemas import DialecticTraceCreate + +# Patterns that indicate the agent abstained from answering +ABSTENTION_PATTERNS = [ + r"don't have (?:enough )?information", + r"cannot (?:find|answer|determine)", + r"no relevant", + r"not (?:enough|sufficient) (?:information|context|data)", + r"unable to (?:find|answer|determine)", + r"no (?:observations|memory|data) (?:found|available)", +] + +_ABSTENTION_REGEX = re.compile("|".join(ABSTENTION_PATTERNS), re.IGNORECASE) + + +def _is_abstention(response: str) -> bool: + """Check if a response indicates abstention from answering.""" + return bool(_ABSTENTION_REGEX.search(response)) + + +async def create_dialectic_trace( + db: AsyncSession, + trace: DialecticTraceCreate, +) -> models.DialecticTrace: + """ + Create a new DialecticTrace record. + + Args: + db: Database session + trace: DialecticTraceCreate schema with trace data + + Returns: + The created DialecticTrace model instance + """ + db_trace = models.DialecticTrace( + workspace_name=trace.workspace_name, + session_name=trace.session_name, + observer=trace.observer, + observed=trace.observed, + query=trace.query, + retrieved_doc_ids=trace.retrieved_doc_ids, + tool_calls=trace.tool_calls, + response=trace.response, + reasoning_level=trace.reasoning_level, + total_duration_ms=trace.total_duration_ms, + input_tokens=trace.input_tokens, + output_tokens=trace.output_tokens, + ) + db.add(db_trace) + await db.flush() + return db_trace + + +async def get_dialectic_traces( + db: AsyncSession, + workspace_name: str, + limit: int = 100, + offset: int = 0, +) -> list[models.DialecticTrace]: + """ + Get dialectic traces for a workspace. + + Args: + db: Database session + workspace_name: Workspace to query traces for + limit: Maximum number of traces to return (default 100) + offset: Number of traces to skip (default 0) + + Returns: + List of DialecticTrace records, ordered by created_at descending + """ + stmt = ( + select(models.DialecticTrace) + .where(models.DialecticTrace.workspace_name == workspace_name) + .order_by(models.DialecticTrace.created_at.desc()) + .limit(limit) + .offset(offset) + ) + result = await db.execute(stmt) + return list(result.scalars().all()) + + +async def get_dialectic_trace_stats( + db: AsyncSession, + workspace_name: str, + since: datetime.datetime | None = None, +) -> dict[str, Any]: + """ + Get aggregate statistics for dialectic traces. + + Args: + db: Database session + workspace_name: Workspace to query stats for + since: Optional datetime to filter traces created after this time + + Returns: + Dictionary with: + - total_queries: Total number of dialectic queries + - avg_duration_ms: Average duration in milliseconds + - abstention_count: Number of queries where agent abstained + - abstention_rate: Ratio of abstentions to total queries + """ + # Build base query + base_filter = models.DialecticTrace.workspace_name == workspace_name + if since is not None: + base_filter = base_filter & (models.DialecticTrace.created_at >= since) + + # Get aggregate stats + stmt = select( + func.count(models.DialecticTrace.id).label("total_queries"), + func.avg(models.DialecticTrace.total_duration_ms).label("avg_duration_ms"), + ).where(base_filter) + + result = await db.execute(stmt) + row = result.one() + total_queries = row.total_queries or 0 + avg_duration_ms = float(row.avg_duration_ms) if row.avg_duration_ms else 0.0 + + # Get all responses to check for abstentions + # (We need to do this in Python since abstention detection uses regex) + responses_stmt = select(models.DialecticTrace.response).where(base_filter) + responses_result = await db.execute(responses_stmt) + responses = [r[0] for r in responses_result.all()] + + abstention_count = sum(1 for r in responses if _is_abstention(r)) + abstention_rate = abstention_count / total_queries if total_queries > 0 else 0.0 + + return { + "total_queries": total_queries, + "avg_duration_ms": avg_duration_ms, + "abstention_count": abstention_count, + "abstention_rate": abstention_rate, + } diff --git a/src/dialectic/core.py b/src/dialectic/core.py index 8a3943ac..df139315 100644 --- a/src/dialectic/core.py +++ b/src/dialectic/core.py @@ -6,6 +6,7 @@ and synthesize responses to queries about a peer. """ import logging +import re import time import uuid from collections.abc import AsyncIterator, Callable @@ -16,6 +17,7 @@ from sqlalchemy.ext.asyncio import AsyncSession from src import crud from src.config import ReasoningLevel, settings from src.dialectic import prompts +from src.schemas import DialecticTraceCreate from src.telemetry import prometheus_metrics from src.telemetry.events import DialecticCompletedEvent, emit from src.telemetry.logging import ( @@ -39,6 +41,33 @@ from src.utils.formatting import format_new_turn_with_timestamp logger = logging.getLogger(__name__) +# Regex to extract document IDs from tool result messages +# Matches patterns like [id:abc123] in formatted observation output +_DOC_ID_PATTERN = re.compile(r"\[id:([a-zA-Z0-9_-]+)\]") + + +def _extract_doc_ids_from_messages(messages: list[dict[str, str]]) -> list[str]: + """ + Extract document IDs from tool_result messages. + + Tool results contain formatted strings like: + [id:abc123] [2025-01-01] The user likes coffee + + Args: + messages: List of conversation messages + + Returns: + List of unique document IDs found in tool results + """ + doc_ids: set[str] = set() + for msg in messages: + if msg.get("role") == "user": + # Tool results appear as user messages in the conversation + content = msg.get("content", "") + matches = _DOC_ID_PATTERN.findall(content) + doc_ids.update(matches) + return list(doc_ids) + class DialecticAgent: """ @@ -287,17 +316,18 @@ class DialecticAgent: return tool_executor, task_name, run_id, start_time - def _log_response_metrics( + async def _log_response_metrics( self, task_name: str, run_id: str | None, start_time: float, + query: str, response_content: str, input_tokens: int, output_tokens: int, cache_read_input_tokens: int | None, cache_creation_input_tokens: int | None, - tool_calls_count: int, + tool_calls_made: list[dict[str, Any]], thinking_content: str | None, iterations: int, ) -> None: @@ -308,15 +338,17 @@ class DialecticAgent: task_name: Metrics task identifier run_id: Run identifier (None if using caller-provided metric_key) start_time: Start time from time.perf_counter() + query: The original query string response_content: The full response text input_tokens: Input token count (actual from API) output_tokens: Output token count (actual from API) cache_read_input_tokens: Cache read tokens (if any) cache_creation_input_tokens: Cache creation tokens (if any) - tool_calls_count: Number of tool calls made + tool_calls_made: List of tool calls made during the response thinking_content: Thinking trace content (if any) iterations: Number of iterations in the tool execution loop """ + tool_calls_count = len(tool_calls_made) accumulate_metric(task_name, "tool_calls", tool_calls_count, "count") if thinking_content: @@ -371,6 +403,28 @@ class DialecticAgent: ) ) + # Persist dialectic trace for meta-cognitive analysis + retrieved_doc_ids = _extract_doc_ids_from_messages(self.messages) + trace = DialecticTraceCreate( + workspace_name=self.workspace_name, + session_name=self.session_name, + observer=self.observer, + observed=self.observed, + query=query, + retrieved_doc_ids=retrieved_doc_ids, + tool_calls=tool_calls_made, + response=response_content, + reasoning_level=self.reasoning_level, + total_duration_ms=elapsed_ms, + input_tokens=input_tokens, + output_tokens=output_tokens, + ) + try: + await crud.create_dialectic_trace(self.db, trace) + except Exception as e: + # Don't fail the request if trace persistence fails + logger.warning(f"Failed to persist dialectic trace: {e}") + async def answer(self, query: str) -> str: """ Answer a query about the peer using agentic tool calling. @@ -419,16 +473,17 @@ class DialecticAgent: trace_name="dialectic_chat", ) - self._log_response_metrics( + await self._log_response_metrics( task_name=task_name, run_id=run_id, start_time=start_time, + query=query, response_content=response.content, input_tokens=response.input_tokens, output_tokens=response.output_tokens, cache_read_input_tokens=response.cache_read_input_tokens, cache_creation_input_tokens=response.cache_creation_input_tokens, - tool_calls_count=len(response.tool_calls_made), + tool_calls_made=response.tool_calls_made, thinking_content=response.thinking_content, iterations=response.iterations, ) @@ -494,16 +549,17 @@ class DialecticAgent: accumulated_content.append(chunk.content) yield chunk.content - self._log_response_metrics( + await self._log_response_metrics( task_name=task_name, run_id=run_id, start_time=start_time, + query=query, response_content="".join(accumulated_content), input_tokens=response.input_tokens, output_tokens=response.output_tokens, cache_read_input_tokens=response.cache_read_input_tokens, cache_creation_input_tokens=response.cache_creation_input_tokens, - tool_calls_count=len(response.tool_calls_made), + tool_calls_made=response.tool_calls_made, thinking_content=response.thinking_content, iterations=response.iterations, ) diff --git a/src/models.py b/src/models.py index c2aa84dd..520006e8 100644 --- a/src/models.py +++ b/src/models.py @@ -11,6 +11,7 @@ from sqlalchemy import ( CheckConstraint, Column, DateTime, + Float, ForeignKey, ForeignKeyConstraint, Identity, @@ -574,3 +575,33 @@ class SessionPeer(Base): internal_metadata: Mapped[dict[str, Any]] joined_at: Mapped[datetime.datetime] left_at: Mapped[datetime.datetime | None] + + +@final +class DialecticTrace(Base): + """Internal logging of dialectic interactions for meta-cognitive analysis.""" + + __tablename__: str = "dialectic_traces" + + id: Mapped[str] = mapped_column(TEXT, default=generate_nanoid, primary_key=True) + workspace_name: Mapped[str] = mapped_column( + ForeignKey("workspaces.name"), nullable=False, index=True + ) + session_name: Mapped[str | None] = mapped_column(TEXT, nullable=True, index=True) + observer: Mapped[str] = mapped_column(TEXT, nullable=False, index=True) + observed: Mapped[str] = mapped_column(TEXT, nullable=False, index=True) + query: Mapped[str] = mapped_column(TEXT, nullable=False) + retrieved_doc_ids: Mapped[list[str]] = mapped_column( + JSONB, default=list, server_default=text("'[]'::jsonb") + ) + tool_calls: Mapped[list[dict[str, Any]]] = mapped_column( + JSONB, default=list, server_default=text("'[]'::jsonb") + ) + response: Mapped[str] = mapped_column(TEXT, nullable=False) + reasoning_level: Mapped[str] = mapped_column(TEXT, nullable=False) + total_duration_ms: Mapped[float] = mapped_column(Float, nullable=False) + input_tokens: Mapped[int] = mapped_column(Integer, nullable=False) + output_tokens: Mapped[int] = mapped_column(Integer, nullable=False) + created_at: Mapped[datetime.datetime] = mapped_column( + DateTime(timezone=True), server_default=func.now(), index=True + ) diff --git a/src/schemas.py b/src/schemas.py index a476ee92..2656032a 100644 --- a/src/schemas.py +++ b/src/schemas.py @@ -829,3 +829,20 @@ class WebhookEndpoint(WebhookEndpointBase): created_at: datetime.datetime model_config = ConfigDict(from_attributes=True, populate_by_name=True) # pyright: ignore + + +class DialecticTraceCreate(BaseModel): + """Internal schema for creating DialecticTrace records. Not exposed via API.""" + + workspace_name: str + session_name: str | None = None + observer: str + observed: str + query: str + retrieved_doc_ids: list[str] = Field(default_factory=list) + tool_calls: list[dict[str, Any]] = Field(default_factory=list) + response: str + reasoning_level: str + total_duration_ms: float + input_tokens: int + output_tokens: int diff --git a/tests/test_dialectic_trace.py b/tests/test_dialectic_trace.py new file mode 100644 index 00000000..df45bafc --- /dev/null +++ b/tests/test_dialectic_trace.py @@ -0,0 +1,376 @@ +"""Tests for DialecticTrace CRUD operations.""" + +import datetime + +import pytest +from nanoid import generate as generate_nanoid +from sqlalchemy.ext.asyncio import AsyncSession + +from src import crud, models +from src.crud.dialectic_trace import _is_abstention +from src.schemas import DialecticTraceCreate + + +class TestDialecticTraceUnit: + """Unit tests for dialectic trace CRUD operations.""" + + @pytest.mark.asyncio + async def test_create_dialectic_trace( + self, + db_session: AsyncSession, + sample_data: tuple[models.Workspace, models.Peer], + ): + """Test creating a dialectic trace record.""" + workspace, peer = sample_data + + trace_data = DialecticTraceCreate( + workspace_name=workspace.name, + session_name=None, + observer=peer.name, + observed=peer.name, + query="What does the user like?", + retrieved_doc_ids=["doc1", "doc2", "doc3"], + tool_calls=[ + {"name": "search_memory", "input": {"query": "likes"}, "id": "call_1"} + ], + response="The user likes coffee and hiking.", + reasoning_level="low", + total_duration_ms=1234.56, + input_tokens=500, + output_tokens=100, + ) + + trace = await crud.create_dialectic_trace(db_session, trace_data) + + assert trace.id is not None + assert len(trace.id) == 21 # nanoid default length + assert trace.workspace_name == workspace.name + assert trace.session_name is None + assert trace.observer == peer.name + assert trace.observed == peer.name + assert trace.query == "What does the user like?" + assert trace.retrieved_doc_ids == ["doc1", "doc2", "doc3"] + assert len(trace.tool_calls) == 1 + assert trace.tool_calls[0]["name"] == "search_memory" + assert trace.response == "The user likes coffee and hiking." + assert trace.reasoning_level == "low" + assert trace.total_duration_ms == 1234.56 + assert trace.input_tokens == 500 + assert trace.output_tokens == 100 + assert trace.created_at is not None + + @pytest.mark.asyncio + async def test_create_dialectic_trace_with_session( + self, + db_session: AsyncSession, + sample_data: tuple[models.Workspace, models.Peer], + ): + """Test creating a trace with a session name.""" + workspace, peer = sample_data + + # Create a session + session = models.Session(name=generate_nanoid(), workspace_name=workspace.name) + db_session.add(session) + await db_session.flush() + + trace_data = DialecticTraceCreate( + workspace_name=workspace.name, + session_name=session.name, + observer=peer.name, + observed=peer.name, + query="Test query", + response="Test response", + reasoning_level="medium", + total_duration_ms=500.0, + input_tokens=200, + output_tokens=50, + ) + + trace = await crud.create_dialectic_trace(db_session, trace_data) + + assert trace.session_name == session.name + + @pytest.mark.asyncio + async def test_get_dialectic_traces( + self, + db_session: AsyncSession, + sample_data: tuple[models.Workspace, models.Peer], + ): + """Test retrieving dialectic traces for a workspace.""" + workspace, peer = sample_data + + # Create multiple traces + for i in range(5): + trace_data = DialecticTraceCreate( + workspace_name=workspace.name, + observer=peer.name, + observed=peer.name, + query=f"Query {i}", + response=f"Response {i}", + reasoning_level="low", + total_duration_ms=100.0 * (i + 1), + input_tokens=100, + output_tokens=50, + ) + await crud.create_dialectic_trace(db_session, trace_data) + + traces = await crud.get_dialectic_traces(db_session, workspace.name) + + assert len(traces) == 5 + # Should be ordered by created_at descending (most recent first) + # Due to fast creation, we check that all traces are present + queries = {t.query for t in traces} + assert queries == {"Query 0", "Query 1", "Query 2", "Query 3", "Query 4"} + + @pytest.mark.asyncio + async def test_get_dialectic_traces_with_limit_offset( + self, + db_session: AsyncSession, + sample_data: tuple[models.Workspace, models.Peer], + ): + """Test pagination of dialectic traces.""" + workspace, peer = sample_data + + # Create 10 traces + for i in range(10): + trace_data = DialecticTraceCreate( + workspace_name=workspace.name, + observer=peer.name, + observed=peer.name, + query=f"Query {i}", + response=f"Response {i}", + reasoning_level="low", + total_duration_ms=100.0, + input_tokens=100, + output_tokens=50, + ) + await crud.create_dialectic_trace(db_session, trace_data) + + # Test limit + traces = await crud.get_dialectic_traces(db_session, workspace.name, limit=3) + assert len(traces) == 3 + + # Test offset + traces = await crud.get_dialectic_traces( + db_session, workspace.name, limit=5, offset=5 + ) + assert len(traces) == 5 + + @pytest.mark.asyncio + async def test_get_dialectic_trace_stats( + self, + db_session: AsyncSession, + sample_data: tuple[models.Workspace, models.Peer], + ): + """Test getting aggregate statistics for dialectic traces.""" + workspace, peer = sample_data + + # Create traces with varying durations + durations = [100.0, 200.0, 300.0, 400.0, 500.0] + for i, duration in enumerate(durations): + trace_data = DialecticTraceCreate( + workspace_name=workspace.name, + observer=peer.name, + observed=peer.name, + query=f"Query {i}", + response=f"Response {i}", + reasoning_level="low", + total_duration_ms=duration, + input_tokens=100, + output_tokens=50, + ) + await crud.create_dialectic_trace(db_session, trace_data) + + stats = await crud.get_dialectic_trace_stats(db_session, workspace.name) + + assert stats["total_queries"] == 5 + assert stats["avg_duration_ms"] == 300.0 # (100+200+300+400+500) / 5 + assert stats["abstention_count"] == 0 + assert stats["abstention_rate"] == 0.0 + + @pytest.mark.asyncio + async def test_get_dialectic_trace_stats_with_abstentions( + self, + db_session: AsyncSession, + sample_data: tuple[models.Workspace, models.Peer], + ): + """Test that abstentions are correctly detected and counted.""" + workspace, peer = sample_data + + responses = [ + "The user likes coffee.", # Not abstention + "I don't have enough information to answer that.", # Abstention + "Based on the observations, the user prefers tea.", # Not abstention + "I cannot find any relevant data about this topic.", # Abstention + "No relevant observations found for this query.", # Abstention + ] + + for i, response in enumerate(responses): + trace_data = DialecticTraceCreate( + workspace_name=workspace.name, + observer=peer.name, + observed=peer.name, + query=f"Query {i}", + response=response, + reasoning_level="low", + total_duration_ms=100.0, + input_tokens=100, + output_tokens=50, + ) + await crud.create_dialectic_trace(db_session, trace_data) + + stats = await crud.get_dialectic_trace_stats(db_session, workspace.name) + + assert stats["total_queries"] == 5 + assert stats["abstention_count"] == 3 + assert stats["abstention_rate"] == 0.6 + + @pytest.mark.asyncio + async def test_get_dialectic_trace_stats_with_since_filter( + self, + db_session: AsyncSession, + sample_data: tuple[models.Workspace, models.Peer], + ): + """Test filtering stats by time.""" + workspace, peer = sample_data + + # Create traces + for i in range(3): + trace_data = DialecticTraceCreate( + workspace_name=workspace.name, + observer=peer.name, + observed=peer.name, + query=f"Query {i}", + response=f"Response {i}", + reasoning_level="low", + total_duration_ms=100.0, + input_tokens=100, + output_tokens=50, + ) + await crud.create_dialectic_trace(db_session, trace_data) + + # Get stats since far future (should return zero) + future_time = datetime.datetime.now(datetime.UTC) + datetime.timedelta(days=1) + stats = await crud.get_dialectic_trace_stats( + db_session, workspace.name, since=future_time + ) + + assert stats["total_queries"] == 0 + assert stats["abstention_count"] == 0 + assert stats["abstention_rate"] == 0.0 + + @pytest.mark.asyncio + async def test_get_dialectic_trace_stats_empty_workspace( + self, + db_session: AsyncSession, + sample_data: tuple[models.Workspace, models.Peer], + ): + """Test stats for workspace with no traces.""" + workspace, _ = sample_data + + stats = await crud.get_dialectic_trace_stats(db_session, workspace.name) + + assert stats["total_queries"] == 0 + assert stats["avg_duration_ms"] == 0.0 + assert stats["abstention_count"] == 0 + assert stats["abstention_rate"] == 0.0 + + +class TestAbstentionDetection: + """Test the abstention detection helper function.""" + + def test_abstention_patterns(self): + """Test that various abstention patterns are detected.""" + abstention_responses = [ + "I don't have information about that.", + "I don't have enough information to answer.", + "I cannot answer this question.", + "I cannot find any relevant data.", + "There is no relevant information available.", + "Not enough information to determine.", + "Not sufficient context available.", + "Unable to find any observations.", + "No observations found for this query.", + "No memory available about this topic.", + "No data found regarding this question.", + ] + + for response in abstention_responses: + assert _is_abstention(response), f"Should detect abstention: {response}" + + def test_non_abstention_responses(self): + """Test that normal responses are not flagged as abstentions.""" + normal_responses = [ + "The user likes coffee.", + "Based on observations, they prefer morning meetings.", + "The data shows a preference for Python.", + "They have mentioned enjoying hiking.", + "According to recent conversations, they work remotely.", + ] + + for response in normal_responses: + assert not _is_abstention(response), ( + f"Should not detect abstention: {response}" + ) + + +class TestDocIdExtraction: + """Test document ID extraction from messages.""" + + def test_extract_doc_ids_from_tool_results(self): + """Test extracting document IDs from formatted tool results.""" + from src.dialectic.core import _extract_doc_ids_from_messages + + messages = [ + {"role": "system", "content": "You are a helpful assistant."}, + {"role": "user", "content": "Query: What does the user like?"}, + { + "role": "assistant", + "content": "Let me search for relevant information.", + }, + { + "role": "user", + "content": """Found 3 observations: + +## Explicit Observations +[id:abc123] [2025-01-01] The user likes coffee +[id:def456] [2025-01-02] The user works remotely + +## Deductive Observations +[id:ghi789] [2025-01-03] The user likely prefers morning meetings""", + }, + {"role": "assistant", "content": "Based on the observations..."}, + ] + + doc_ids = _extract_doc_ids_from_messages(messages) + + assert set(doc_ids) == {"abc123", "def456", "ghi789"} + + def test_extract_doc_ids_no_matches(self): + """Test extraction when no document IDs are present.""" + from src.dialectic.core import _extract_doc_ids_from_messages + + messages = [ + {"role": "system", "content": "You are a helpful assistant."}, + {"role": "user", "content": "Hello!"}, + {"role": "assistant", "content": "Hi there!"}, + ] + + doc_ids = _extract_doc_ids_from_messages(messages) + + assert doc_ids == [] + + def test_extract_doc_ids_duplicates_removed(self): + """Test that duplicate IDs are deduplicated.""" + from src.dialectic.core import _extract_doc_ids_from_messages + + messages = [ + { + "role": "user", + "content": "[id:abc123] First mention\n[id:abc123] Same ID again", + }, + ] + + doc_ids = _extract_doc_ids_from_messages(messages) + + assert doc_ids == ["abc123"]