Fix SSE with query auth token param.

This commit is contained in:
febux 2026-09-02 22:23:16 +03:00
parent 6ae7edc8f5
commit 3c80d21f0a
6 changed files with 118 additions and 18 deletions

View file

@ -1455,7 +1455,7 @@ Health/metrics exempt from auth.
| POST | `/api/v1/documents` | user JWT | multipart → reserve credit → MinIO put → row `queued` → publish `DocumentUploaded``202 {document_id, correlation_id}`. `402` if no credit. `400` bad mime/size. |
| GET | `/api/v1/documents/{id}` | user JWT | status + stage + filename (for polling UI) |
| GET | `/api/v1/reports/{document_id}` | user JWT | `202 {status, stage}` while not done; `200 {markdown, findings, ...}` when done |
| GET | `/api/v1/reports/{document_id}/events` | user JWT | SSE-стрим статусов: безымянные кадры с payload `analysisResultSchema` (camelCase, `pending/processing/completed/failed`); терминальные `completed/failed` закрывают стрим; служебные события `timeout`/`error`; keep-alive комментарии между кадрами |
| GET | `/api/v1/reports/{document_id}/events` | user JWT (header или `?access_token=` — для нативного EventSource) | SSE-стрим статусов: безымянные кадры с payload `analysisResultSchema` (camelCase, `pending/processing/completed/failed`); терминальные `completed/failed` закрывают стрим; служебные события `timeout`/`error`; keep-alive комментарии между кадрами |
| GET | `/api/v1/me` | user JWT | `{telegram_id, credits_left}` |
### B2B endpoints (`X-API-Key`, `api/routes/b2b.py`)

View file

@ -6,7 +6,7 @@ from collections.abc import AsyncIterator
from typing import Annotated, Any
from uuid import UUID
from fastapi import Depends, Header, HTTPException, Request
from fastapi import Depends, Header, HTTPException, Query, Request
from sqlalchemy import select
from sqlalchemy.ext.asyncio import AsyncSession
@ -339,18 +339,8 @@ class CurrentUser:
self.telegram_id = telegram_id
async def require_current_user(
session: AsyncSessionDep,
authorization: Annotated[str | None, Header()] = None,
) -> CurrentUser:
"""Validate `Authorization: Bearer <user-jwt>` and return the user.
This is the common auth gate for bot, web, and Mini App users.
"""
if not authorization or not authorization.lower().startswith("bearer "):
raise HTTPException(status_code=401, detail="Missing or invalid Authorization header")
token = authorization[7:].strip()
async def resolve_user_from_token(session: AsyncSession, token: str) -> CurrentUser:
"""Verify a raw access JWT and return the user (shared auth core)."""
try:
claims = verify_access_token(token)
except TokenExpiredError as exc:
@ -368,9 +358,57 @@ async def require_current_user(
return CurrentUser(user_id=claims.sub, telegram_id=claims.telegram_id)
async def require_current_user(
session: AsyncSessionDep,
authorization: Annotated[str | None, Header()] = None,
) -> CurrentUser:
"""Validate `Authorization: Bearer <user-jwt>` and return the user.
This is the common auth gate for bot, web, and Mini App users.
"""
if not authorization or not authorization.lower().startswith("bearer "):
raise HTTPException(status_code=401, detail="Missing or invalid Authorization header")
token = authorization[7:].strip()
return await resolve_user_from_token(session, token)
CurrentUserDep = Annotated[CurrentUser, Depends(require_current_user)]
async def require_current_user_header_or_query(
session: AsyncSessionDep,
authorization: Annotated[str | None, Header()] = None,
access_token: Annotated[str | None, Query()] = None,
) -> CurrentUser:
"""Auth for EventSource-compatible endpoints (SSE).
Native `EventSource` cannot set request headers, so besides the usual
`Authorization: Bearer` header this also accepts the access JWT via the
`access_token` query parameter:
new EventSource(`/api/v1/reports/{id}/events?access_token=<jwt>`)
Caveat: query strings may end up in proxy access logs pass short-lived
access tokens only, never refresh tokens.
"""
token: str | None = None
if authorization and authorization.lower().startswith("bearer "):
token = authorization[7:].strip()
elif access_token:
token = access_token.strip()
if not token:
raise HTTPException(
status_code=401,
detail="Missing auth: use Authorization header or access_token query param",
)
return await resolve_user_from_token(session, token)
EventSourceUserDep = Annotated[CurrentUser, Depends(require_current_user_header_or_query)]
async def get_credits(session: AsyncSession, user_id: UUID) -> int:
credits = await UserRepository(session).get_credits(user_id)
return credits if credits is not None else 0

View file

@ -552,6 +552,17 @@ SSE-стрим статусов анализа (`text/event-stream`). Аутен
сервер сам поллит БД (`SSE_POLL_INTERVAL_SECONDS`) и шлёт событие при
смене статуса.
**Auth:** `Authorization: Bearer <jwt>` **или** query-параметр
`?access_token=<jwt>` — нативный `EventSource` не умеет ставить заголовки,
поэтому для web-SPA:
```js
new EventSource(`/api/v1/reports/${docId}/events?access_token=${jwt}`)
```
Нюанс: query-строка может попасть в access-логи прокси — передавать только
короткоживущие access-токены (не refresh).
Каждый кадр данных (безымянное событие → `onmessage` в `EventSource`)
несёт payload контракта web-клиента `analysisResultSchema` (camelCase;
см. [`api/schemas/analysis_result.py`](../schemas/analysis_result.py)):

View file

@ -15,6 +15,7 @@ from sqlalchemy.ext.asyncio import AsyncSession, async_sessionmaker
from src.contract_check.api.deps import (
AsyncSessionDep,
CurrentUserDep,
EventSourceUserDep,
fetch_document_status_for_user,
)
from src.contract_check.api.schemas import ReportInProgressResponse, ReportResponse
@ -78,11 +79,13 @@ async def stream_report_events(
request: Request,
session: AsyncSessionDep,
document_id: UUID,
user: CurrentUserDep,
user: EventSourceUserDep,
) -> StreamingResponse:
"""Server-Sent Events stream of the document analysis status.
Auth and ownership are checked before the stream starts (401/404 as
Auth accepts either `Authorization: Bearer <jwt>` or, for native
`EventSource` (which cannot set headers), the `access_token` query
parameter. Ownership is checked before the stream starts (401/404 as
regular HTTP errors). Every data frame carries an `AnalysisResult`
payload (the web client's `analysisResultSchema` — parse each event's
`data` with it):

View file

@ -52,6 +52,11 @@ def _set_env() -> None:
"JWT_ALGORITHM": "HS256",
"JWT_ACCESS_TTL_MINUTES": "1440",
"TELEGRAM_BOT_TOKEN": "it-test-bot-token:it-test-secret",
# httpx ASGITransport buffers the whole response — infinite SSE streams
# can only complete via the max-stream deadline (which emits the
# `timeout` event). Keep it short so streaming tests finish fast.
"SSE_POLL_INTERVAL_SECONDS": "0.25",
"SSE_MAX_STREAM_SECONDS": "5",
}
for k, v in env_vars.items():
os.environ[k] = v

View file

@ -122,6 +122,46 @@ async def test_sse_requires_auth(client: httpx.AsyncClient) -> None:
assert response.status_code == 401
async def test_sse_accepts_access_token_query_param(
client: httpx.AsyncClient,
db_session, # noqa: ANN001
infra: dict[str, str],
pdf_bytes: bytes,
) -> None:
"""Native EventSource cannot send headers — auth must work via query param."""
telegram_id = 999100106
await _create_user_with_credits(db_session, telegram_id)
token = await user_token(client, infra, telegram_id)
document_id = await _upload(client, token, pdf_bytes)
events: list[tuple[str, dict[str, Any]]] = []
async with client.stream(
"GET",
f"/api/v1/reports/{document_id}/events",
params={"access_token": token},
) as response:
assert response.status_code == 200
buffer = ""
async for chunk in response.aiter_text():
buffer += chunk
events = _parse_events(buffer)
if events and events[-1][0] in ("timeout", "error"):
break
assert len(events) >= 1
name, payload = events[0]
assert name == "message"
assert payload["id"] == document_id
async def test_sse_rejects_garbage_access_token(client: httpx.AsyncClient) -> None:
response = await client.get(
f"/api/v1/reports/{uuid.uuid4()}/events",
params={"access_token": "not-a-jwt"},
)
assert response.status_code == 401
async def test_sse_unknown_document_returns_404(
client: httpx.AsyncClient,
infra: dict[str, str],
@ -285,5 +325,8 @@ async def test_sse_emits_keepalive_comments_while_unchanged(
assert "data: " in raw
assert raw.count(": keep-alive") >= 2
# Only one data frame (the initial snapshot) while nothing changes.
assert len(_parse_events(raw)) == 1
# At least the initial data frame; when the stream hits the short test
# deadline it also emits a `timeout` event and closes.
events = _parse_events(raw)
assert len(events) >= 1
assert events[0][0] == "message"