dp-zp-agent/app/routes.py

510 lines
10 KiB
Python

from __future__ import annotations
import asyncio
import json
import os
from typing import Any
from fastapi import (
APIRouter,
Depends,
Header,
HTTPException,
Request,
status,
)
from fastapi.responses import JSONResponse
from pydantic import BaseModel, Field
from app.security import (
expected_gitea_repository,
repository_name_from_payload,
require_search_api_key,
require_sync_api_key,
same_repository,
validate_secret,
verify_gitea_signature,
webhook_should_pull_git,
)
from scripts.common import (
DB_FILE,
ZPWIKI_ROOT,
)
from scripts.rag_utils import build_rag_context
from scripts.rebuild_index import (
ReindexInProgressError,
rebuild_index,
)
from scripts.search_utils import search_database
router = APIRouter()
class SearchRequest(BaseModel):
query: str = Field(
...,
min_length=1,
max_length=500,
)
limit: int = Field(
default=10,
ge=1,
le=50,
)
published_only: bool = False
max_per_document: int = Field(
default=1,
ge=0,
le=10,
)
class RagRequest(BaseModel):
query: str = Field(
...,
min_length=1,
max_length=500,
description=(
"Otázka alebo vyhľadávací dotaz "
"používateľa nad ZP Wiki."
),
)
limit: int = Field(
default=5,
ge=1,
le=20,
description=(
"Maximálny počet relevantných "
"zdrojov pre RAG kontext."
),
)
published_only: bool = Field(
default=False,
description=(
"Ak je true, použijú sa iba "
"publikované dokumenty."
),
)
max_per_document: int = Field(
default=1,
ge=0,
le=10,
description=(
"Maximálny počet chunkov z jedného "
"dokumentu. Hodnota 1 preferuje "
"rôzne dokumenty."
),
)
class SyncRequest(BaseModel):
pull_git: bool = Field(
default=False,
description=(
"Pred reindexovaním vykoná "
"git pull --ff-only."
),
)
@router.get(
"/health",
include_in_schema=False,
)
def health() -> dict[str, Any]:
return {
"status": "ok",
"database_exists": (
DB_FILE.exists()
),
"database_path": str(
DB_FILE
),
"search_engine": (
"hybrid_fts5_embeddings"
),
"rag_enabled": True,
"zpwiki_root": str(
ZPWIKI_ROOT
),
"zpwiki_exists": (
ZPWIKI_ROOT.exists()
),
"security_configured": all(
bool(
os.getenv(
name,
"",
).strip()
)
for name in (
"WEBHOOK_SECRET",
"SYNC_API_KEY",
"SEARCH_API_KEY",
"EXPECTED_GITEA_REPOSITORY",
)
),
}
@router.post(
"/rag",
operation_id=(
"retrieve_zpwiki_context"
),
summary=(
"Vyhľadaj informácie v ZP Wiki"
),
description=(
"Použi tento nástroj pri otázkach "
"o ZP Wiki, študentoch, autoroch, "
"záverečných prácach, témach, rokoch, "
"projektoch alebo dokumentoch. "
"Nástroj vykoná hybridné FTS5 a "
"embeddingové vyhľadávanie a pripraví "
"zdrojovo podložený RAG kontext."
),
dependencies=[
Depends(
require_search_api_key
)
],
)
def rag(
request: RagRequest,
) -> dict[str, Any]:
try:
response = build_rag_context(
DB_FILE,
request.query,
limit=request.limit,
published_only=(
request.published_only
),
max_per_document=(
request.max_per_document
),
)
except FileNotFoundError as error:
raise HTTPException(
status_code=500,
detail=str(
error
),
) from error
except ValueError as error:
raise HTTPException(
status_code=400,
detail=str(
error
),
) from error
except RuntimeError as error:
raise HTTPException(
status_code=500,
detail=str(
error
),
) from error
return response
@router.post(
"/search",
dependencies=[
Depends(
require_search_api_key
)
],
include_in_schema=False,
)
def search(
request: SearchRequest,
) -> dict[str, Any]:
try:
response = search_database(
DB_FILE,
request.query,
request.limit,
published_only=(
request.published_only
),
max_per_document=(
request.max_per_document
),
)
except FileNotFoundError as error:
raise HTTPException(
status_code=500,
detail=str(
error
),
) from error
except ValueError as error:
raise HTTPException(
status_code=400,
detail=str(
error
),
) from error
except RuntimeError as error:
raise HTTPException(
status_code=500,
detail=str(
error
),
) from error
results = response[
"results"
]
return {
"query": request.query,
"engine": response[
"engine"
],
"strategies": response[
"strategies"
],
"count": len(
results
),
"results": results,
}
@router.post(
"/sync",
dependencies=[
Depends(
require_sync_api_key
)
],
include_in_schema=False,
)
def sync(
request: SyncRequest,
) -> dict[str, Any]:
try:
result = rebuild_index(
pull_git=(
request.pull_git
)
)
except ReindexInProgressError as error:
raise HTTPException(
status_code=409,
detail=str(
error
),
) from error
except RuntimeError as error:
raise HTTPException(
status_code=500,
detail=str(
error
),
) from error
return {
"status": "ok",
"pull_git": (
request.pull_git
),
"duration_seconds": (
result[
"duration_seconds"
]
),
"counts": result[
"counts"
],
}
@router.post(
"/webhook/gitea",
response_model=None,
include_in_schema=False,
)
async def gitea_webhook(
request: Request,
x_gitea_event: str | None = Header(
default=None,
alias="X-Gitea-Event",
),
x_gitea_signature: str | None = Header(
default=None,
alias="X-Gitea-Signature",
),
) -> dict[str, Any] | JSONResponse:
raw_body = await request.body()
secret = validate_secret(
"WEBHOOK_SECRET"
)
if not verify_gitea_signature(
raw_body,
x_gitea_signature,
secret,
):
raise HTTPException(
status_code=(
status.HTTP_401_UNAUTHORIZED
),
detail=(
"Neplatný webhook podpis"
),
)
try:
payload = json.loads(
raw_body.decode(
"utf-8"
)
)
except (
UnicodeDecodeError,
json.JSONDecodeError,
) as error:
raise HTTPException(
status_code=400,
detail=(
"Webhook payload nie je "
"platný JSON"
),
) from error
if not isinstance(
payload,
dict,
):
raise HTTPException(
status_code=400,
detail=(
"Webhook payload musí byť "
"JSON objekt"
),
)
if not x_gitea_event:
raise HTTPException(
status_code=400,
detail=(
"Chýba hlavička "
"X-Gitea-Event"
),
)
if (
x_gitea_event.casefold()
!= "push"
):
return JSONResponse(
status_code=(
status.HTTP_202_ACCEPTED
),
content={
"status": "ignored",
"reason": (
"unsupported_event"
),
"event": (
x_gitea_event
),
},
)
repository_name = (
repository_name_from_payload(
payload
)
)
if repository_name is None:
raise HTTPException(
status_code=400,
detail=(
"Webhook payload neobsahuje "
"repository.full_name"
),
)
expected_repository = (
expected_gitea_repository()
)
if not same_repository(
repository_name,
expected_repository,
):
raise HTTPException(
status_code=403,
detail=(
"Webhook patrí neočakávanému "
"repozitáru"
),
)
try:
result = await asyncio.to_thread(
rebuild_index,
pull_git=(
webhook_should_pull_git()
),
)
except ReindexInProgressError as error:
raise HTTPException(
status_code=409,
detail=str(
error
),
) from error
except RuntimeError as error:
raise HTTPException(
status_code=500,
detail=str(
error
),
) from error
return {
"status": "ok",
"event": (
x_gitea_event
),
"repository": (
repository_name
),
"verified_by": (
"hmac_sha256"
),
"duration_seconds": (
result[
"duration_seconds"
]
),
"counts": result[
"counts"
],
}