|
| 1 | +from fastapi import APIRouter, HTTPException |
| 2 | +from pydantic import BaseModel |
| 3 | +from api.services.llm_explainer import TransactionExplainer |
| 4 | +from api.services.llm_query import QueryTranslator |
| 5 | +from api.services.llm_context import MultiModalContextHandler |
| 6 | +from api.services.llm_validation import ResponseValidator |
| 7 | +from typing import List, Dict, Any |
| 8 | + |
| 9 | +router = APIRouter(prefix="/api/v1/llm", tags=["llm"]) |
| 10 | +explainer = TransactionExplainer() |
| 11 | +query_translator = QueryTranslator() |
| 12 | +context_handler = MultiModalContextHandler() |
| 13 | +validator = ResponseValidator() |
| 14 | + |
| 15 | + |
| 16 | + |
| 17 | + |
| 18 | +class ExplainRequest(BaseModel): |
| 19 | + tx_details: str |
| 20 | + |
| 21 | +class ExplainResponse(BaseModel): |
| 22 | + explanation: str |
| 23 | + |
| 24 | +@router.post("/explain", response_model=ExplainResponse) |
| 25 | +async def explain_transaction(request: ExplainRequest): |
| 26 | + try: |
| 27 | + explanation = await explainer.explain(request.tx_details) |
| 28 | + return ExplainResponse(explanation=explanation) |
| 29 | + except Exception as e: |
| 30 | + raise HTTPException(status_code=500, detail=str(e)) |
| 31 | + |
| 32 | +class QueryRequest(BaseModel): |
| 33 | + query: str |
| 34 | + |
| 35 | +class QueryResponse(BaseModel): |
| 36 | + sql: str |
| 37 | + |
| 38 | +@router.post("/query", response_model=QueryResponse) |
| 39 | +async def translate_query(request: QueryRequest): |
| 40 | + try: |
| 41 | + sql = query_translator.translate_to_sql(request.query) |
| 42 | + return QueryResponse(sql=sql) |
| 43 | + except ValueError as e: |
| 44 | + raise HTTPException(status_code=400, detail=str(e)) |
| 45 | + except Exception as e: |
| 46 | + raise HTTPException(status_code=500, detail=str(e)) |
| 47 | + |
| 48 | +class ContextRequest(BaseModel): |
| 49 | + edges: List[Dict[str, Any]] = [] |
| 50 | + data_points: List[float] = [] |
| 51 | + |
| 52 | +class ContextResponse(BaseModel): |
| 53 | + graph_summary: str |
| 54 | + time_series_trend: str |
| 55 | + mermaid: str |
| 56 | + |
| 57 | +@router.post("/context", response_model=ContextResponse) |
| 58 | +async def get_multimodal_context(request: ContextRequest): |
| 59 | + try: |
| 60 | + summary = context_handler.serialize_and_summarize_graph(request.edges) |
| 61 | + trend = context_handler.extract_time_series(request.data_points) |
| 62 | + mermaid = context_handler.generate_mermaid_diagram([], request.edges) |
| 63 | + return ContextResponse( |
| 64 | + graph_summary=summary, |
| 65 | + time_series_trend=trend, |
| 66 | + mermaid=mermaid |
| 67 | + ) |
| 68 | + except Exception as e: |
| 69 | + raise HTTPException(status_code=500, detail=str(e)) |
| 70 | + |
| 71 | +class ValidateRequest(BaseModel): |
| 72 | + raw_response: Dict[str, Any] |
| 73 | + context: str |
| 74 | + |
| 75 | +class ValidateResponse(BaseModel): |
| 76 | + validated_response: Dict[str, Any] |
| 77 | + |
| 78 | +@router.post("/validate", response_model=ValidateResponse) |
| 79 | +async def validate_response(request: ValidateRequest): |
| 80 | + try: |
| 81 | + validated = validator.validate_and_guard(request.raw_response, request.context) |
| 82 | + return ValidateResponse(validated_response=validated) |
| 83 | + except ValueError as e: |
| 84 | + raise HTTPException(status_code=400, detail=str(e)) |
| 85 | + except Exception as e: |
| 86 | + raise HTTPException(status_code=500, detail=str(e)) |
| 87 | + |
0 commit comments