Jai-rathore29's picture
Deploy: DocAgent backend (deterministic date-anomaly fix)
f65e025
Raw
History Blame Contribute Delete
5.24 kB
"""Agent tools — the capabilities the orchestrator can call.
Each tool is a thin async wrapper over a service, plus a `ToolSpec`
(JSON-schema description) the LLM sees. Tools are document-scoped: the
orchestrator passes the active document id in via a closure context.
"""
from __future__ import annotations
import json
from dataclasses import dataclass
from typing import Any, Awaitable, Callable
from app.core.logging import get_logger
from app.llm.base import ToolSpec
from app.services import classification, extraction, qa, storage, summary
from app.services import anomaly as anomaly_svc
from app.services import vectorstore
log = get_logger(__name__)
@dataclass
class ToolContext:
document_id: str | None
provider: str | None = None
# collected so the API can surface citations from the last query
citations: list = None # type: ignore
def __post_init__(self):
if self.citations is None:
self.citations = []
ToolFn = Callable[[ToolContext, dict[str, Any]], Awaitable[str]]
# -------------------------- tool implementations --------------------------
async def _query_document(ctx: ToolContext, args: dict) -> str:
if not ctx.document_id:
return "No document is currently loaded."
question = args.get("question", "")
res = await qa.answer(ctx.document_id, question, provider=ctx.provider)
ctx.citations.extend(res.citations)
cites = "\n".join(
f"[{i+1}] (page {c.page+1}) {c.text[:160]}" for i, c in enumerate(res.citations)
)
return f"{res.answer}\n\nGrounding passages:\n{cites}"
async def _classify(ctx: ToolContext, args: dict) -> str:
detail = storage.get(ctx.document_id) if ctx.document_id else None
if not detail:
return "No document loaded."
if detail.classification:
c = detail.classification
return f"Document type: {c.doc_type} (confidence {c.confidence:.0%}). {c.rationale}"
c = await classification.classify(detail.markdown or "", ctx.provider)
detail.classification = c
storage.save(detail)
return f"Document type: {c.doc_type} (confidence {c.confidence:.0%}). {c.rationale}"
async def _get_extracted_data(ctx: ToolContext, args: dict) -> str:
detail = storage.get(ctx.document_id) if ctx.document_id else None
if not detail or not detail.extraction:
return "No extracted data yet."
fields = {f.name: f.value for f in detail.extraction.fields}
return json.dumps({
"schema": detail.extraction.schema_name,
"fields": fields,
"tables": [t.title or "table" for t in detail.extraction.tables],
"entities": [{"type": e.type, "value": e.value} for e in detail.extraction.entities],
}, default=str)
async def _summarize(ctx: ToolContext, args: dict) -> str:
detail = storage.get(ctx.document_id) if ctx.document_id else None
if not detail:
return "No document loaded."
if detail.summary:
return detail.summary
s = await summary.summarize(detail.markdown or "", ctx.provider)
detail.summary = s
storage.save(detail)
return s
async def _flag_anomalies(ctx: ToolContext, args: dict) -> str:
detail = storage.get(ctx.document_id) if ctx.document_id else None
if not detail or not detail.extraction:
return "No extracted data to check."
doc_type = detail.classification.doc_type if detail.classification else None
items = await anomaly_svc.detect(detail.markdown or "", detail.extraction,
doc_type, ctx.provider)
if not items:
return "No anomalies detected. The document looks consistent."
return "\n".join(f"- [{a.severity.value}] {a.field or ''}: {a.message}" for a in items)
# ----------------------------- registry -----------------------------------
TOOL_SPECS: list[ToolSpec] = [
ToolSpec(
name="query_document",
description="Retrieve grounded passages from the document and answer a question about its contents.",
parameters={
"type": "object",
"properties": {"question": {"type": "string", "description": "The question to answer"}},
"required": ["question"],
},
),
ToolSpec(
name="classify_document",
description="Get the document type/category with confidence.",
parameters={"type": "object", "properties": {}},
),
ToolSpec(
name="get_extracted_data",
description="Return the structured fields, tables, and entities already extracted from the document.",
parameters={"type": "object", "properties": {}},
),
ToolSpec(
name="summarize_document",
description="Produce or fetch a concise summary of the document.",
parameters={"type": "object", "properties": {}},
),
ToolSpec(
name="flag_anomalies",
description="Check for missing required fields, low-confidence values, and inconsistencies.",
parameters={"type": "object", "properties": {}},
),
]
TOOL_FNS: dict[str, ToolFn] = {
"query_document": _query_document,
"classify_document": _classify,
"get_extracted_data": _get_extracted_data,
"summarize_document": _summarize,
"flag_anomalies": _flag_anomalies,
}