main upravy a testy upravy

This commit is contained in:
Ján Pták 2026-08-14 22:24:45 +02:00
parent 2871bc853a
commit 0a3d898b2c
5 changed files with 1401 additions and 949 deletions

View File

@ -1,32 +1,13 @@
from __future__ import annotations from __future__ import annotations
import asyncio import asyncio
import hashlib
import hmac
import json
import os
import sys import sys
from contextlib import asynccontextmanager from contextlib import asynccontextmanager
from pathlib import Path from pathlib import Path
from typing import Any
from fastapi import ( from fastapi import FastAPI
Depends,
FastAPI,
Header,
HTTPException,
Request,
Security,
status,
)
from fastapi.middleware.cors import CORSMiddleware from fastapi.middleware.cors import CORSMiddleware
from fastapi.responses import JSONResponse
from fastapi.security import (
APIKeyHeader,
HTTPAuthorizationCredentials,
HTTPBearer,
)
from pydantic import BaseModel, Field
PROJECT_ROOT = Path(__file__).resolve().parents[1] PROJECT_ROOT = Path(__file__).resolve().parents[1]
@ -36,203 +17,44 @@ if str(PROJECT_ROOT) not in sys.path:
str(PROJECT_ROOT), str(PROJECT_ROOT),
) )
from scripts.common import (
DB_FILE, from app.routes import (
ZPWIKI_ROOT, RagRequest,
SearchRequest,
SyncRequest,
gitea_webhook,
health,
rag,
router,
search,
sync,
)
from app.security import (
MIN_SECRET_LENGTH,
SEARCH_API_KEY_HEADER,
SYNC_API_KEY_HEADER,
expected_gitea_repository,
repository_name_from_payload,
require_search_api_key,
require_sync_api_key,
required_environment_value,
same_repository,
search_api_key_scheme,
search_bearer_scheme,
sync_api_key_scheme,
validate_secret,
validate_security_configuration,
verify_gitea_signature,
webhook_should_pull_git,
) )
from scripts.embedding_utils import embed_query from scripts.embedding_utils import embed_query
from scripts.rag_utils import build_rag_context
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"
OPENWEBUI_ORIGIN = ( OPENWEBUI_ORIGIN = (
"https://ui.tukekemt.xyz" "https://ui.tukekemt.xyz"
) )
search_api_key_scheme = APIKeyHeader(
name=SEARCH_API_KEY_HEADER,
scheme_name="SearchApiKey",
auto_error=False,
description=(
"API kľúč pre vyhľadávanie v ZP Wiki "
"cez hlavičku X-API-Key."
),
)
search_bearer_scheme = HTTPBearer(
scheme_name="SearchBearer",
auto_error=False,
description=(
"Bearer token pre vyhľadávanie v ZP Wiki. "
"Používa hodnotu SEARCH_API_KEY."
),
)
sync_api_key_scheme = APIKeyHeader(
name=SYNC_API_KEY_HEADER,
scheme_name="SyncApiKey",
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=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."
),
)
def required_environment_value(
name: str,
) -> str:
value = os.getenv(
name,
"",
).strip()
if not value:
raise RuntimeError(
"Chýba povinná environment "
f"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ň "
f"{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 @asynccontextmanager
async def lifespan( async def lifespan(
_: FastAPI, _: FastAPI,
@ -277,577 +99,7 @@ app.add_middleware(
) )
def require_search_api_key( app.include_router(
api_key: str | None = Security( router
search_api_key_scheme
),
bearer: (
HTTPAuthorizationCredentials | None
) = Security(
search_bearer_scheme
),
) -> None:
expected = validate_secret(
"SEARCH_API_KEY"
)
supplied_credentials: list[str] = []
if api_key:
supplied_credentials.append(
api_key
)
if (
bearer is not None
and bearer.scheme.casefold()
== "bearer"
and bearer.credentials
):
supplied_credentials.append(
bearer.credentials
)
valid = any(
hmac.compare_digest(
supplied,
expected,
)
for supplied
in supplied_credentials
)
if not valid:
raise HTTPException(
status_code=(
status.HTTP_401_UNAUTHORIZED
),
detail=(
"Neplatný alebo chýbajúci "
"API kľúč"
),
headers={
"WWW-Authenticate": (
"Bearer"
),
},
)
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()
)
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",
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",
)
),
}
@app.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
@app.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,
}
@app.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"
],
}
@app.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"
],
}

509
app/routes.py Normal file
View File

@ -0,0 +1,509 @@
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"
],
}

303
app/security.py Normal file
View File

@ -0,0 +1,303 @@
from __future__ import annotations
import hashlib
import hmac
import os
from typing import Any
from fastapi import HTTPException, Security, status
from fastapi.security import (
APIKeyHeader,
HTTPAuthorizationCredentials,
HTTPBearer,
)
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,
scheme_name="SearchApiKey",
auto_error=False,
description=(
"API kľúč pre vyhľadávanie v ZP Wiki "
"cez hlavičku X-API-Key."
),
)
search_bearer_scheme = HTTPBearer(
scheme_name="SearchBearer",
auto_error=False,
description=(
"Bearer token pre vyhľadávanie v ZP Wiki. "
"Používa hodnotu SEARCH_API_KEY."
),
)
sync_api_key_scheme = APIKeyHeader(
name=SYNC_API_KEY_HEADER,
scheme_name="SyncApiKey",
auto_error=False,
description=(
"API kľúč pre manuálne spustenie "
"reindexovania."
),
)
def required_environment_value(
name: str,
) -> str:
value = os.getenv(
name,
"",
).strip()
if not value:
raise RuntimeError(
"Chýba povinná environment "
f"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ň "
f"{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()
def require_search_api_key(
api_key: str | None = Security(
search_api_key_scheme
),
bearer: (
HTTPAuthorizationCredentials | None
) = Security(
search_bearer_scheme
),
) -> None:
expected = validate_secret(
"SEARCH_API_KEY"
)
supplied_credentials: list[str] = []
if api_key:
supplied_credentials.append(
api_key
)
if (
bearer is not None
and bearer.scheme.casefold()
== "bearer"
and bearer.credentials
):
supplied_credentials.append(
bearer.credentials
)
valid = any(
hmac.compare_digest(
supplied,
expected,
)
for supplied
in supplied_credentials
)
if not valid:
raise HTTPException(
status_code=(
status.HTTP_401_UNAUTHORIZED
),
detail=(
"Neplatný alebo chýbajúci "
"API kľúč"
),
headers={
"WWW-Authenticate": (
"Bearer"
),
},
)
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()
)
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(),
)

View File

@ -8,7 +8,10 @@ import pytest
from fastapi.testclient import TestClient from fastapi.testclient import TestClient
import app.main as main import app.main as main
from scripts.rebuild_index import ReindexInProgressError import app.routes as routes
from scripts.rebuild_index import (
ReindexInProgressError,
)
WEBHOOK_SECRET = "w" * 64 WEBHOOK_SECRET = "w" * 64
@ -29,9 +32,13 @@ def fake_rebuild_result() -> dict:
} }
def sign(body: bytes) -> str: def sign(
body: bytes,
) -> str:
return hmac.new( return hmac.new(
WEBHOOK_SECRET.encode("utf-8"), WEBHOOK_SECRET.encode(
"utf-8"
),
body, body,
hashlib.sha256, hashlib.sha256,
).hexdigest() ).hexdigest()
@ -39,20 +46,26 @@ def sign(body: bytes) -> str:
@pytest.fixture @pytest.fixture
def client() -> TestClient: def client() -> TestClient:
with TestClient(main.app) as test_client: with TestClient(
main.app
) as test_client:
yield test_client yield test_client
def test_startup_rejects_missing_secret( def test_startup_rejects_missing_secret(
monkeypatch: pytest.MonkeyPatch, monkeypatch: pytest.MonkeyPatch,
) -> None: ) -> None:
monkeypatch.delenv("WEBHOOK_SECRET") monkeypatch.delenv(
"WEBHOOK_SECRET"
)
with pytest.raises( with pytest.raises(
RuntimeError, RuntimeError,
match="WEBHOOK_SECRET", match="WEBHOOK_SECRET",
): ):
with TestClient(main.app): with TestClient(
main.app
):
pass pass
@ -68,22 +81,42 @@ def test_startup_rejects_short_secret(
RuntimeError, RuntimeError,
match="aspoň 32", match="aspoň 32",
): ):
with TestClient(main.app): with TestClient(
main.app
):
pass pass
def test_health_endpoint( def test_health_endpoint(
client: TestClient, client: TestClient,
) -> None: ) -> None:
response = client.get("/health") response = client.get(
"/health"
)
assert response.status_code == 200 assert (
response.status_code
== 200
)
payload = response.json() payload = response.json()
assert payload["status"] == "ok" assert (
assert payload["search_engine"] == "hybrid_fts5_embeddings" payload["status"]
assert payload["security_configured"] is True == "ok"
)
assert (
payload["search_engine"]
== "hybrid_fts5_embeddings"
)
assert (
payload[
"security_configured"
]
is True
)
def test_search_endpoint_uses_shared_search_logic( def test_search_endpoint_uses_shared_search_logic(
@ -91,14 +124,20 @@ def test_search_endpoint_uses_shared_search_logic(
monkeypatch: pytest.MonkeyPatch, monkeypatch: pytest.MonkeyPatch,
) -> None: ) -> None:
monkeypatch.setattr( monkeypatch.setattr(
main, routes,
"search_database", "search_database",
lambda *args, **kwargs: { lambda *args, **kwargs: {
"engine": "sqlite_fts5", "engine": (
"strategies": ["all_terms"], "sqlite_fts5"
),
"strategies": [
"all_terms"
],
"results": [ "results": [
{ {
"chunk_id": "test::0", "chunk_id": (
"test::0"
),
"published": True, "published": True,
} }
], ],
@ -108,7 +147,9 @@ def test_search_endpoint_uses_shared_search_logic(
response = client.post( response = client.post(
"/search", "/search",
headers={ headers={
"X-API-Key": SEARCH_API_KEY, "X-API-Key": (
SEARCH_API_KEY
),
}, },
json={ json={
"query": "jan ptak", "query": "jan ptak",
@ -116,10 +157,24 @@ def test_search_endpoint_uses_shared_search_logic(
}, },
) )
assert response.status_code == 200
assert response.json()["count"] == 1
assert ( assert (
response.json()["results"][0]["chunk_id"] response.status_code
== 200
)
assert (
response.json()[
"count"
]
== 1
)
assert (
response.json()[
"results"
][0][
"chunk_id"
]
== "test::0" == "test::0"
) )
@ -130,14 +185,19 @@ def test_search_rejects_empty_query(
response = client.post( response = client.post(
"/search", "/search",
headers={ headers={
"X-API-Key": SEARCH_API_KEY, "X-API-Key": (
SEARCH_API_KEY
),
}, },
json={ json={
"query": "", "query": "",
}, },
) )
assert response.status_code == 422 assert (
response.status_code
== 422
)
def test_search_rejects_missing_api_key( def test_search_rejects_missing_api_key(
@ -151,7 +211,10 @@ def test_search_rejects_missing_api_key(
}, },
) )
assert response.status_code == 401 assert (
response.status_code
== 401
)
def test_search_rejects_wrong_api_key( def test_search_rejects_wrong_api_key(
@ -160,7 +223,9 @@ def test_search_rejects_wrong_api_key(
response = client.post( response = client.post(
"/search", "/search",
headers={ headers={
"X-API-Key": "x" * 64, "X-API-Key": (
"x" * 64
),
}, },
json={ json={
"query": "jan ptak", "query": "jan ptak",
@ -168,7 +233,10 @@ def test_search_rejects_wrong_api_key(
}, },
) )
assert response.status_code == 401 assert (
response.status_code
== 401
)
def test_sync_rejects_missing_api_key( def test_sync_rejects_missing_api_key(
@ -181,7 +249,10 @@ def test_sync_rejects_missing_api_key(
}, },
) )
assert response.status_code == 401 assert (
response.status_code
== 401
)
def test_sync_rejects_wrong_api_key( def test_sync_rejects_wrong_api_key(
@ -190,14 +261,19 @@ def test_sync_rejects_wrong_api_key(
response = client.post( response = client.post(
"/sync", "/sync",
headers={ headers={
"X-API-Key": "x" * 64, "X-API-Key": (
"x" * 64
),
}, },
json={ json={
"pull_git": False, "pull_git": False,
}, },
) )
assert response.status_code == 401 assert (
response.status_code
== 401
)
def test_sync_accepts_valid_api_key( def test_sync_accepts_valid_api_key(
@ -205,36 +281,54 @@ def test_sync_accepts_valid_api_key(
monkeypatch: pytest.MonkeyPatch, monkeypatch: pytest.MonkeyPatch,
) -> None: ) -> None:
monkeypatch.setattr( monkeypatch.setattr(
main, routes,
"rebuild_index", "rebuild_index",
lambda pull_git=False: fake_rebuild_result(), lambda pull_git=False: (
fake_rebuild_result()
),
) )
response = client.post( response = client.post(
"/sync", "/sync",
headers={ headers={
"X-API-Key": SYNC_API_KEY, "X-API-Key": (
SYNC_API_KEY
),
}, },
json={ json={
"pull_git": False, "pull_git": False,
}, },
) )
assert response.status_code == 200 assert (
assert response.json()["status"] == "ok" response.status_code
== 200
)
assert (
response.json()[
"status"
]
== "ok"
)
def test_sync_returns_conflict_when_reindex_is_running( def test_sync_returns_conflict_when_reindex_is_running(
client: TestClient, client: TestClient,
monkeypatch: pytest.MonkeyPatch, monkeypatch: pytest.MonkeyPatch,
) -> None: ) -> None:
def busy(*args, **kwargs): def busy(
raise ReindexInProgressError( *args,
"Reindexovanie už prebieha" **kwargs,
):
raise (
ReindexInProgressError(
"Reindexovanie už prebieha"
)
) )
monkeypatch.setattr( monkeypatch.setattr(
main, routes,
"rebuild_index", "rebuild_index",
busy, busy,
) )
@ -242,14 +336,19 @@ def test_sync_returns_conflict_when_reindex_is_running(
response = client.post( response = client.post(
"/sync", "/sync",
headers={ headers={
"X-API-Key": SYNC_API_KEY, "X-API-Key": (
SYNC_API_KEY
),
}, },
json={ json={
"pull_git": False, "pull_git": False,
}, },
) )
assert response.status_code == 409 assert (
response.status_code
== 409
)
def test_webhook_rejects_invalid_signature( def test_webhook_rejects_invalid_signature(
@ -258,22 +357,35 @@ def test_webhook_rejects_invalid_signature(
body = json.dumps( body = json.dumps(
{ {
"repository": { "repository": {
"full_name": "KEMT/zpwiki", "full_name": (
"KEMT/zpwiki"
),
} }
} }
).encode("utf-8") ).encode(
"utf-8"
)
response = client.post( response = client.post(
"/webhook/gitea", "/webhook/gitea",
content=body, content=body,
headers={ headers={
"Content-Type": "application/json", "Content-Type": (
"X-Gitea-Event": "push", "application/json"
"X-Gitea-Signature": "0" * 64, ),
"X-Gitea-Event": (
"push"
),
"X-Gitea-Signature": (
"0" * 64
),
}, },
) )
assert response.status_code == 401 assert (
response.status_code
== 401
)
def test_webhook_rejects_invalid_json_with_valid_signature( def test_webhook_rejects_invalid_json_with_valid_signature(
@ -285,13 +397,24 @@ def test_webhook_rejects_invalid_json_with_valid_signature(
"/webhook/gitea", "/webhook/gitea",
content=body, content=body,
headers={ headers={
"Content-Type": "application/json", "Content-Type": (
"X-Gitea-Event": "push", "application/json"
"X-Gitea-Signature": sign(body), ),
"X-Gitea-Event": (
"push"
),
"X-Gitea-Signature": (
sign(
body
)
),
}, },
) )
assert response.status_code == 400 assert (
response.status_code
== 400
)
def test_webhook_requires_event_header( def test_webhook_requires_event_header(
@ -300,21 +423,34 @@ def test_webhook_requires_event_header(
body = json.dumps( body = json.dumps(
{ {
"repository": { "repository": {
"full_name": "KEMT/zpwiki", "full_name": (
"KEMT/zpwiki"
),
} }
} }
).encode("utf-8") ).encode(
"utf-8"
)
response = client.post( response = client.post(
"/webhook/gitea", "/webhook/gitea",
content=body, content=body,
headers={ headers={
"Content-Type": "application/json", "Content-Type": (
"X-Gitea-Signature": sign(body), "application/json"
),
"X-Gitea-Signature": (
sign(
body
)
),
}, },
) )
assert response.status_code == 400 assert (
response.status_code
== 400
)
def test_webhook_ignores_non_push_event( def test_webhook_ignores_non_push_event(
@ -323,13 +459,20 @@ def test_webhook_ignores_non_push_event(
) -> None: ) -> None:
calls = 0 calls = 0
def fake_rebuild(*args, **kwargs): def fake_rebuild(
*args,
**kwargs,
):
nonlocal calls nonlocal calls
calls += 1 calls += 1
return fake_rebuild_result()
return (
fake_rebuild_result()
)
monkeypatch.setattr( monkeypatch.setattr(
main, routes,
"rebuild_index", "rebuild_index",
fake_rebuild, fake_rebuild,
) )
@ -337,23 +480,45 @@ def test_webhook_ignores_non_push_event(
body = json.dumps( body = json.dumps(
{ {
"repository": { "repository": {
"full_name": "KEMT/zpwiki", "full_name": (
"KEMT/zpwiki"
),
} }
} }
).encode("utf-8") ).encode(
"utf-8"
)
response = client.post( response = client.post(
"/webhook/gitea", "/webhook/gitea",
content=body, content=body,
headers={ headers={
"Content-Type": "application/json", "Content-Type": (
"X-Gitea-Event": "issues", "application/json"
"X-Gitea-Signature": sign(body), ),
"X-Gitea-Event": (
"issues"
),
"X-Gitea-Signature": (
sign(
body
)
),
}, },
) )
assert response.status_code == 202 assert (
assert response.json()["status"] == "ignored" response.status_code
== 202
)
assert (
response.json()[
"status"
]
== "ignored"
)
assert calls == 0 assert calls == 0
@ -363,22 +528,37 @@ def test_webhook_rejects_unexpected_repository(
body = json.dumps( body = json.dumps(
{ {
"repository": { "repository": {
"full_name": "OTHER/repository", "full_name": (
"OTHER/repository"
),
} }
} }
).encode("utf-8") ).encode(
"utf-8"
)
response = client.post( response = client.post(
"/webhook/gitea", "/webhook/gitea",
content=body, content=body,
headers={ headers={
"Content-Type": "application/json", "Content-Type": (
"X-Gitea-Event": "push", "application/json"
"X-Gitea-Signature": sign(body), ),
"X-Gitea-Event": (
"push"
),
"X-Gitea-Signature": (
sign(
body
)
),
}, },
) )
assert response.status_code == 403 assert (
response.status_code
== 403
)
def test_webhook_accepts_signed_push( def test_webhook_accepts_signed_push(
@ -386,36 +566,59 @@ def test_webhook_accepts_signed_push(
monkeypatch: pytest.MonkeyPatch, monkeypatch: pytest.MonkeyPatch,
) -> None: ) -> None:
monkeypatch.setattr( monkeypatch.setattr(
main, routes,
"rebuild_index", "rebuild_index",
lambda pull_git=False: fake_rebuild_result(), lambda pull_git=False: (
fake_rebuild_result()
),
) )
body = json.dumps( body = json.dumps(
{ {
"repository": { "repository": {
"full_name": "KEMT/zpwiki", "full_name": (
"KEMT/zpwiki"
),
} }
} }
).encode("utf-8") ).encode(
"utf-8"
)
response = client.post( response = client.post(
"/webhook/gitea", "/webhook/gitea",
content=body, content=body,
headers={ headers={
"Content-Type": "application/json", "Content-Type": (
"X-Gitea-Event": "push", "application/json"
"X-Gitea-Signature": sign(body), ),
"X-Gitea-Event": (
"push"
),
"X-Gitea-Signature": (
sign(
body
)
),
}, },
) )
assert response.status_code == 200
assert ( assert (
response.json()["verified_by"] response.status_code
== 200
)
assert (
response.json()[
"verified_by"
]
== "hmac_sha256" == "hmac_sha256"
) )
assert ( assert (
response.json()["repository"] response.json()[
"repository"
]
== "KEMT/zpwiki" == "KEMT/zpwiki"
) )
@ -424,13 +627,18 @@ def test_webhook_returns_conflict_when_reindex_is_running(
client: TestClient, client: TestClient,
monkeypatch: pytest.MonkeyPatch, monkeypatch: pytest.MonkeyPatch,
) -> None: ) -> None:
def busy(*args, **kwargs): def busy(
raise ReindexInProgressError( *args,
"Reindexovanie už prebieha" **kwargs,
):
raise (
ReindexInProgressError(
"Reindexovanie už prebieha"
)
) )
monkeypatch.setattr( monkeypatch.setattr(
main, routes,
"rebuild_index", "rebuild_index",
busy, busy,
) )
@ -438,19 +646,34 @@ def test_webhook_returns_conflict_when_reindex_is_running(
body = json.dumps( body = json.dumps(
{ {
"repository": { "repository": {
"full_name": "KEMT/zpwiki", "full_name": (
"KEMT/zpwiki"
),
} }
} }
).encode("utf-8") ).encode(
"utf-8"
)
response = client.post( response = client.post(
"/webhook/gitea", "/webhook/gitea",
content=body, content=body,
headers={ headers={
"Content-Type": "application/json", "Content-Type": (
"X-Gitea-Event": "push", "application/json"
"X-Gitea-Signature": sign(body), ),
"X-Gitea-Event": (
"push"
),
"X-Gitea-Signature": (
sign(
body
)
),
}, },
) )
assert response.status_code == 409 assert (
response.status_code
== 409
)

View File

@ -7,6 +7,7 @@ import pytest
from fastapi.testclient import TestClient from fastapi.testclient import TestClient
import app.main as main_module import app.main as main_module
import app.routes as routes
import scripts.rag_utils as rag_utils import scripts.rag_utils as rag_utils
from scripts.rag_utils import ( from scripts.rag_utils import (
ANSWER_FORMAT, ANSWER_FORMAT,
@ -19,6 +20,7 @@ from scripts.rag_utils import (
SEARCH_API_KEY = "a" * 64 SEARCH_API_KEY = "a" * 64
@pytest.fixture @pytest.fixture
def client( def client(
security_environment, security_environment,
@ -27,6 +29,7 @@ def client(
main_module.app main_module.app
) )
def sample_result() -> dict[str, Any]: def sample_result() -> dict[str, Any]:
return { return {
"chunk_id": ( "chunk_id": (
@ -74,29 +77,56 @@ def test_build_source() -> None:
1, 1,
) )
assert source["source_id"] == "S1" assert (
source["source_id"]
== "S1"
)
assert source["title"] == "Ján Holp" assert (
assert source["author"] == "Daniel Hladek" source["title"]
== "Ján Holp"
)
assert source["source_url"] == ( assert (
source["author"]
== "Daniel Hladek"
)
assert source[
"source_url"
] == (
"https://zp.kemt.fei.tuke.sk/" "https://zp.kemt.fei.tuke.sk/"
"students/2016/jan_holp" "students/2016/jan_holp"
) )
assert source["published"] is True assert (
source["published"]
is True
)
assert source["retrieval"] == { assert source[
"match_strategy": "any_term", "retrieval"
] == {
"match_strategy": (
"any_term"
),
"fts_rank": 11, "fts_rank": 11,
"vector_rank": 1, "vector_rank": 1,
"vector_score": 0.863072, "vector_score": (
"hybrid_score": 0.02811129, 0.863072
),
"hybrid_score": (
0.02811129
),
} }
# Interná identifikácia zdroja nemá byť # Interná identifikácia zdroja
# používateľská citation hodnota. # nemá byť používateľská
assert "citation" not in source # citation hodnota.
assert (
"citation"
not in source
)
def test_build_context_text() -> None: def test_build_context_text() -> None:
@ -105,12 +135,21 @@ def test_build_context_text() -> None:
1, 1,
) )
context = build_context_text( context = (
[source] build_context_text(
[source]
)
) )
assert "ZDROJ S1" in context assert (
assert "Názov dokumentu: Ján Holp" in context "ZDROJ S1"
in context
)
assert (
"Názov dokumentu: Ján Holp"
in context
)
assert ( assert (
"Autor dokumentu: Daniel Hladek" "Autor dokumentu: Daniel Hladek"
@ -135,7 +174,11 @@ def test_build_context_text() -> None:
def test_build_context_text_empty() -> None: def test_build_context_text_empty() -> None:
context = build_context_text([]) context = (
build_context_text(
[]
)
)
assert ( assert (
"nenašli relevantné zdroje" "nenašli relevantné zdroje"
@ -168,7 +211,10 @@ def test_rag_instructions_require_grounding() -> None:
in instructions in instructions
) )
assert "source_url" in instructions assert (
"source_url"
in instructions
)
def test_answer_format() -> None: def test_answer_format() -> None:
@ -188,14 +234,19 @@ def test_answer_format() -> None:
assert ( assert (
"<source_url>" "<source_url>"
in ANSWER_FORMAT["template"] in ANSWER_FORMAT[
"template"
]
) )
def test_build_rag_context( def test_build_rag_context(
monkeypatch: pytest.MonkeyPatch, monkeypatch: pytest.MonkeyPatch,
) -> None: ) -> None:
captured: dict[str, Any] = {} captured: dict[
str,
Any,
] = {}
def fake_search_database( def fake_search_database(
db_path: Path, db_path: Path,
@ -205,15 +256,25 @@ def test_build_rag_context(
published_only: bool, published_only: bool,
max_per_document: int, max_per_document: int,
) -> dict[str, Any]: ) -> dict[str, Any]:
captured["db_path"] = db_path captured[
captured["query"] = query "db_path"
captured["limit"] = limit ] = db_path
captured["published_only"] = (
published_only captured[
) "query"
captured["max_per_document"] = ( ] = query
max_per_document
) captured[
"limit"
] = limit
captured[
"published_only"
] = published_only
captured[
"max_per_document"
] = max_per_document
return { return {
"engine": ( "engine": (
@ -237,15 +298,17 @@ def test_build_rag_context(
"/tmp/test.sqlite" "/tmp/test.sqlite"
) )
response = build_rag_context( response = (
db_path, build_rag_context(
( db_path,
"V akom roku robil Ján Holp " (
"diplomovú prácu?" "V akom roku robil Ján Holp "
), "diplomovú prácu?"
limit=5, ),
published_only=True, limit=5,
max_per_document=1, published_only=True,
max_per_document=1,
)
) )
assert captured == { assert captured == {
@ -259,28 +322,49 @@ def test_build_rag_context(
"max_per_document": 1, "max_per_document": 1,
} }
assert response["engine"] == ( assert (
"hybrid_fts5_embeddings" response[
"engine"
]
== "hybrid_fts5_embeddings"
) )
assert response["strategies"] == [ assert (
"any_term" response[
] "strategies"
]
assert response["source_count"] == 1 == [
"any_term"
]
)
assert ( assert (
response["sources"][0]["title"] response[
"source_count"
]
== 1
)
assert (
response[
"sources"
][0][
"title"
]
== "Ján Holp" == "Ján Holp"
) )
assert ( assert (
"Diplomová práca 2021" "Diplomová práca 2021"
in response["context"] in response[
"context"
]
) )
assert ( assert (
response["answer_format"][ response[
"answer_format"
][
"internal_source_ids_visible" "internal_source_ids_visible"
] ]
is False is False
@ -312,17 +396,34 @@ def test_build_rag_context_without_results(
fake_search_database, fake_search_database,
) )
response = build_rag_context( response = (
Path("/tmp/test.sqlite"), build_rag_context(
"neexistujúca téma", Path(
"/tmp/test.sqlite"
),
"neexistujúca téma",
)
) )
assert response["source_count"] == 0 assert (
assert response["sources"] == [] response[
"source_count"
]
== 0
)
assert (
response[
"sources"
]
== []
)
assert ( assert (
"nenašli relevantné zdroje" "nenašli relevantné zdroje"
in response["context"] in response[
"context"
]
) )
@ -347,12 +448,17 @@ def test_rag_endpoint(
), ),
"context": ( "context": (
"ZDROJ S1\n" "ZDROJ S1\n"
"Názov dokumentu: Ján Holp" "Názov dokumentu: "
"Ján Holp"
), ),
"sources": [ "sources": [
{ {
"source_id": "S1", "source_id": (
"title": "Ján Holp", "S1"
),
"title": (
"Ján Holp"
),
"source_url": ( "source_url": (
"https://example.test/" "https://example.test/"
"jan_holp" "jan_holp"
@ -369,15 +475,30 @@ def test_rag_endpoint(
published_only: bool, published_only: bool,
max_per_document: int, max_per_document: int,
) -> dict[str, Any]: ) -> dict[str, Any]:
assert query == "Ján Holp" assert (
assert limit == 5 query
assert published_only is False == "Ján Holp"
assert max_per_document == 1 )
assert (
limit
== 5
)
assert (
published_only
is False
)
assert (
max_per_document
== 1
)
return expected return expected
monkeypatch.setattr( monkeypatch.setattr(
main_module, routes,
"build_rag_context", "build_rag_context",
fake_build_rag_context, fake_build_rag_context,
) )
@ -390,12 +511,21 @@ def test_rag_endpoint(
), ),
}, },
json={ json={
"query": "Ján Holp", "query": (
"Ján Holp"
),
}, },
) )
assert response.status_code == 200 assert (
assert response.json() == expected response.status_code
== 200
)
assert (
response.json()
== expected
)
def test_rag_endpoint_with_bearer( def test_rag_endpoint_with_bearer(
@ -423,12 +553,14 @@ def test_rag_endpoint_with_bearer(
"answer_format": ( "answer_format": (
ANSWER_FORMAT ANSWER_FORMAT
), ),
"context": "bez výsledkov", "context": (
"bez výsledkov"
),
"sources": [], "sources": [],
} }
monkeypatch.setattr( monkeypatch.setattr(
main_module, routes,
"build_rag_context", "build_rag_context",
fake_build_rag_context, fake_build_rag_context,
) )
@ -437,7 +569,8 @@ def test_rag_endpoint_with_bearer(
"/rag", "/rag",
headers={ headers={
"Authorization": ( "Authorization": (
f"Bearer {SEARCH_API_KEY}" f"Bearer "
f"{SEARCH_API_KEY}"
), ),
}, },
json={ json={
@ -445,7 +578,10 @@ def test_rag_endpoint_with_bearer(
}, },
) )
assert response.status_code == 200 assert (
response.status_code
== 200
)
def test_rag_endpoint_without_api_key( def test_rag_endpoint_without_api_key(
@ -454,11 +590,16 @@ def test_rag_endpoint_without_api_key(
response = client.post( response = client.post(
"/rag", "/rag",
json={ json={
"query": "Ján Holp", "query": (
"Ján Holp"
),
}, },
) )
assert response.status_code == 401 assert (
response.status_code
== 401
)
def test_rag_endpoint_empty_query( def test_rag_endpoint_empty_query(
@ -476,7 +617,10 @@ def test_rag_endpoint_empty_query(
}, },
) )
assert response.status_code == 422 assert (
response.status_code
== 422
)
def test_openapi_exposes_rag_only( def test_openapi_exposes_rag_only(
@ -486,15 +630,36 @@ def test_openapi_exposes_rag_only(
"/openapi.json" "/openapi.json"
) )
assert response.status_code == 200 assert (
response.status_code
== 200
)
paths = response.json()[ paths = response.json()[
"paths" "paths"
] ]
assert "/rag" in paths assert (
"/rag"
in paths
)
assert "/search" not in paths assert (
assert "/sync" not in paths "/search"
assert "/health" not in paths not in paths
assert "/webhook/gitea" not in paths )
assert (
"/sync"
not in paths
)
assert (
"/health"
not in paths
)
assert (
"/webhook/gitea"
not in paths
)