dp-zp-agent/app/main.py

530 lines
11 KiB
Python

from __future__ import annotations
import asyncio
import hashlib
import hmac
import json
import os
import sys
from contextlib import asynccontextmanager
from pathlib import Path
from typing import Any
from fastapi import (
Depends,
FastAPI,
Header,
HTTPException,
Request,
Security,
status,
)
from fastapi.responses import JSONResponse
from fastapi.security import APIKeyHeader
from pydantic import BaseModel, Field
PROJECT_ROOT = Path(__file__).resolve().parents[1]
if str(PROJECT_ROOT) not in sys.path:
sys.path.insert(0, str(PROJECT_ROOT))
from scripts.common import DB_FILE, ZPWIKI_ROOT
from scripts.rebuild_index import (
ReindexInProgressError,
rebuild_index,
)
from scripts.search_utils import search_database
MIN_SECRET_LENGTH = 32
SEARCH_API_KEY_HEADER = "X-API-Key"
SYNC_API_KEY_HEADER = "X-API-Key"
search_api_key_scheme = APIKeyHeader(
name=SEARCH_API_KEY_HEADER,
auto_error=False,
description="API kľúč pre vyhľadávanie v zpwiki.",
)
sync_api_key_scheme = APIKeyHeader(
name=SYNC_API_KEY_HEADER,
auto_error=False,
description="API kľúč pre manuálne spustenie reindexovania.",
)
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=3,
ge=0,
le=10,
)
class SyncRequest(BaseModel):
pull_git: bool = Field(
default=False,
description="Pred reindexovaním vykoná git pull --ff-only.",
)
def required_environment_value(name: str) -> str:
value = os.getenv(name, "").strip()
if not value:
raise RuntimeError(
f"Chýba povinná environment premenná {name}"
)
return value
def validate_secret(name: str) -> str:
value = required_environment_value(name)
if len(value) < MIN_SECRET_LENGTH:
raise RuntimeError(
f"{name} musí mať aspoň {MIN_SECRET_LENGTH} znakov"
)
return value
def expected_gitea_repository() -> str:
value = required_environment_value(
"EXPECTED_GITEA_REPOSITORY"
)
if "/" not in value:
raise RuntimeError(
"EXPECTED_GITEA_REPOSITORY musí mať tvar "
"vlastník/repozitár"
)
return value
def webhook_should_pull_git() -> bool:
value = os.getenv(
"WEBHOOK_PULL_GIT",
"false",
).strip().casefold()
return value in {
"1",
"true",
"yes",
"on",
}
def validate_security_configuration() -> None:
validate_secret("WEBHOOK_SECRET")
validate_secret("SYNC_API_KEY")
validate_secret("SEARCH_API_KEY")
expected_gitea_repository()
@asynccontextmanager
async def lifespan(_: FastAPI):
validate_security_configuration()
yield
app = FastAPI(
title="ZP Agent API",
description=(
"API pre vyhľadávanie v repozitári "
"záverečných prác zpwiki."
),
version="0.6.0",
lifespan=lifespan,
)
def require_search_api_key(
api_key: str | None = Security(
search_api_key_scheme
),
) -> None:
expected = validate_secret(
"SEARCH_API_KEY"
)
if (
not api_key
or not hmac.compare_digest(
api_key,
expected,
)
):
raise HTTPException(
status_code=status.HTTP_401_UNAUTHORIZED,
detail="Neplatný alebo chýbajúci API kľúč",
headers={
"WWW-Authenticate": "ApiKey",
},
)
def require_sync_api_key(
api_key: str | None = Security(
sync_api_key_scheme
),
) -> None:
expected = validate_secret(
"SYNC_API_KEY"
)
if (
not api_key
or not hmac.compare_digest(
api_key,
expected,
)
):
raise HTTPException(
status_code=status.HTTP_401_UNAUTHORIZED,
detail="Neplatný alebo chýbajúci API kľúč",
headers={
"WWW-Authenticate": "ApiKey",
},
)
def verify_gitea_signature(
raw_body: bytes,
signature: str | None,
secret: str,
) -> bool:
if not signature:
return False
supplied = signature.strip().casefold()
# Kompatibilita podpisu.
if supplied.startswith("sha256="):
supplied = supplied.removeprefix(
"sha256="
)
if len(supplied) != 64:
return False
try:
int(supplied, 16)
except ValueError:
return False
expected = hmac.new(
secret.encode("utf-8"),
raw_body,
hashlib.sha256,
).hexdigest()
return hmac.compare_digest(
expected,
supplied,
)
def repository_name_from_payload(
payload: dict[str, Any],
) -> str | None:
repository = payload.get("repository")
if not isinstance(repository, dict):
return None
value = (
repository.get("full_name")
or repository.get("name")
)
if not isinstance(value, str):
return None
value = value.strip()
return value or None
def same_repository(
actual: str,
expected: str,
) -> bool:
return hmac.compare_digest(
actual.casefold(),
expected.casefold(),
)
@app.get("/health")
def health() -> dict[str, Any]:
return {
"status": "ok",
"database_exists": DB_FILE.exists(),
"database_path": str(DB_FILE),
"search_engine": "sqlite_fts5",
"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",
)
),
}
@app.post(
"/search",
dependencies=[
Depends(require_search_api_key)
],
)
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,
}
@app.post(
"/sync",
dependencies=[
Depends(require_sync_api_key)
],
)
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"],
}
@app.post(
"/webhook/gitea",
response_model=None,
)
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"],
}