feat(api): add GraphQL endpoint at /graphql with Strawberry

- Add strawberry-graphql[fastapi] dependency
- Implement GraphQL schema covering documents, pipelines, settings, users
- Enable GraphiQL playground at /graphql
- Mount GraphQL router in main.py
- Add 24 tests for all query types and auth enforcement
- Update docs/API.md with GraphQL documentation section

Co-authored-by: christianlouis <361235+christianlouis@users.noreply.github.com>
This commit is contained in:
copilot-swe-agent[bot]
2026-03-11 22:17:25 +00:00
parent 2b834d1405
commit a41ded535f
5 changed files with 975 additions and 1 deletions
+431
View File
@@ -0,0 +1,431 @@
"""
GraphQL API endpoint for DocuElevate.
Provides a flexible query interface alongside the existing REST API.
Schema covers: documents, pipelines, settings, and users.
Endpoint: /graphql
GraphiQL playground: /graphql (via browser)
"""
from __future__ import annotations
import logging
from datetime import datetime
from typing import Annotated, Any
import strawberry
from fastapi import Depends, Request
from sqlalchemy.orm import Session
from strawberry.fastapi import GraphQLRouter
from app.auth import get_current_user
from app.config import settings
from app.database import get_db
from app.models import ApplicationSettings, FileRecord, Pipeline, PipelineStep, UserProfile
logger = logging.getLogger(__name__)
# ---------------------------------------------------------------------------
# Strawberry types
# ---------------------------------------------------------------------------
@strawberry.type
class DocumentType:
"""A processed document stored in the system."""
id: int
owner_id: str | None
original_filename: str | None
local_filename: str
file_size: int
mime_type: str | None
document_title: str | None
is_duplicate: bool
ocr_quality_score: int | None
pipeline_id: int | None
created_at: datetime | None
@strawberry.type
class PipelineStepType:
"""A single step within a processing pipeline."""
id: int
pipeline_id: int
position: int
step_type: str
label: str | None
enabled: bool
created_at: datetime | None
@strawberry.type
class PipelineType:
"""A processing pipeline with its ordered steps."""
id: int
owner_id: str | None
name: str
description: str | None
is_default: bool
is_active: bool
steps: list[PipelineStepType]
created_at: datetime | None
updated_at: datetime | None
@strawberry.type
class SettingType:
"""An application configuration setting stored in the database."""
id: int
key: str
value: str | None
created_at: datetime | None
updated_at: datetime | None
@strawberry.type
class UserType:
"""A user profile in the system."""
id: int
user_id: str
display_name: str | None
is_blocked: bool
subscription_tier: str | None
onboarding_completed: bool
created_at: datetime | None
# ---------------------------------------------------------------------------
# Conversion helpers
# ---------------------------------------------------------------------------
def _document_from_record(rec: FileRecord) -> DocumentType:
return DocumentType(
id=rec.id,
owner_id=rec.owner_id,
original_filename=rec.original_filename,
local_filename=rec.local_filename,
file_size=rec.file_size,
mime_type=rec.mime_type,
document_title=rec.document_title,
is_duplicate=rec.is_duplicate,
ocr_quality_score=rec.ocr_quality_score,
pipeline_id=rec.pipeline_id,
created_at=rec.created_at,
)
def _pipeline_step_from_record(step: PipelineStep) -> PipelineStepType:
return PipelineStepType(
id=step.id,
pipeline_id=step.pipeline_id,
position=step.position,
step_type=step.step_type,
label=step.label,
enabled=step.enabled,
created_at=step.created_at,
)
def _pipeline_from_record(pipeline: Pipeline, db: Session) -> PipelineType:
steps = db.query(PipelineStep).filter(PipelineStep.pipeline_id == pipeline.id).order_by(PipelineStep.position).all()
return PipelineType(
id=pipeline.id,
owner_id=pipeline.owner_id,
name=pipeline.name,
description=pipeline.description,
is_default=pipeline.is_default,
is_active=pipeline.is_active,
steps=[_pipeline_step_from_record(s) for s in steps],
created_at=pipeline.created_at,
updated_at=pipeline.updated_at,
)
def _setting_from_record(setting: ApplicationSettings) -> SettingType:
return SettingType(
id=setting.id,
key=setting.key,
value=setting.value,
created_at=setting.created_at,
updated_at=setting.updated_at,
)
def _user_from_profile(profile: UserProfile) -> UserType:
return UserType(
id=profile.id,
user_id=profile.user_id,
display_name=profile.display_name,
is_blocked=profile.is_blocked,
subscription_tier=profile.subscription_tier,
onboarding_completed=profile.onboarding_completed,
created_at=profile.created_at,
)
# ---------------------------------------------------------------------------
# Context helpers
# ---------------------------------------------------------------------------
# Keys that contain sensitive data and must never be returned via GraphQL
_SENSITIVE_SETTING_KEYS: frozenset[str] = frozenset(
{
"openai_api_key",
"azure_ai_key",
"session_secret",
"database_url",
"redis_url",
"dropbox_app_secret",
"dropbox_refresh_token",
"google_drive_credentials_json",
"onedrive_client_secret",
"onedrive_refresh_token",
"smtp_password",
"nextcloud_password",
"s3_secret_access_key",
"ftp_password",
"sftp_password",
"webdav_password",
"stripe_secret_key",
"stripe_webhook_secret",
"sentry_dsn",
"social_auth_google_client_secret",
"social_auth_microsoft_client_secret",
"social_auth_apple_private_key",
"social_auth_dropbox_app_secret",
}
)
def _get_current_user_id(user: dict[str, Any] | None) -> str | None:
"""Extract the stable user identifier from the user dict."""
if not user:
return None
return user.get("preferred_username") or user.get("email") or user.get("id") or None
def _get_db_and_user(info: strawberry.types.Info) -> tuple[Session, dict[str, Any] | None]:
"""Extract the database session and current user from the Strawberry context."""
db: Session = info.context["db"]
user: dict[str, Any] | None = info.context.get("user")
return db, user
def _require_auth(user: dict[str, Any] | None) -> None:
"""Raise an error when authentication is enabled and no valid user is present."""
if settings.auth_enabled and not user:
raise strawberry.exceptions.StrawberryGraphQLError("Authentication required")
def _require_admin(user: dict[str, Any] | None) -> None:
"""Raise an error when the current user is not an admin.
When ``auth_enabled`` is *False* (single-user / development mode) all
callers are implicitly treated as administrators.
"""
if not settings.auth_enabled:
# Single-user mode: no auth, treat caller as admin
return
_require_auth(user)
if not (user and user.get("is_admin")):
raise strawberry.exceptions.StrawberryGraphQLError("Admin access required")
# ---------------------------------------------------------------------------
# Query resolvers
# ---------------------------------------------------------------------------
@strawberry.type
class Query:
"""Root query type for the DocuElevate GraphQL API."""
@strawberry.field(description="List documents, optionally filtered by owner.")
def documents(
self,
info: strawberry.types.Info,
owner_id: str | None = None,
limit: int = 20,
offset: int = 0,
) -> list[DocumentType]:
"""Return a paginated list of documents.
When *auth_enabled* the caller must be authenticated. Non-admin users
receive only their own documents; admins may query any *owner_id*.
"""
db, user = _get_db_and_user(info)
_require_auth(user)
limit = max(1, min(limit, 100))
offset = max(0, offset)
query = db.query(FileRecord)
if settings.auth_enabled and user:
is_admin = user.get("is_admin", False)
current_user_id = _get_current_user_id(user)
if not is_admin:
# Non-admins can only see their own documents
query = query.filter(FileRecord.owner_id == current_user_id)
elif owner_id:
query = query.filter(FileRecord.owner_id == owner_id)
elif owner_id:
query = query.filter(FileRecord.owner_id == owner_id)
records = query.order_by(FileRecord.created_at.desc()).offset(offset).limit(limit).all()
return [_document_from_record(r) for r in records]
@strawberry.field(description="Fetch a single document by ID.")
def document(self, info: strawberry.types.Info, id: int) -> DocumentType | None:
"""Return one document by its primary key, or *null* if not found."""
db, user = _get_db_and_user(info)
_require_auth(user)
rec = db.query(FileRecord).filter(FileRecord.id == id).first()
if rec is None:
return None
if settings.auth_enabled and user:
is_admin = user.get("is_admin", False)
current_user_id = _get_current_user_id(user)
if not is_admin and rec.owner_id != current_user_id:
return None
return _document_from_record(rec)
@strawberry.field(description="List processing pipelines.")
def pipelines(
self,
info: strawberry.types.Info,
owner_id: str | None = None,
limit: int = 20,
offset: int = 0,
) -> list[PipelineType]:
"""Return a paginated list of pipelines."""
db, user = _get_db_and_user(info)
_require_auth(user)
limit = max(1, min(limit, 100))
offset = max(0, offset)
query = db.query(Pipeline)
if settings.auth_enabled and user:
is_admin = user.get("is_admin", False)
current_user_id = _get_current_user_id(user)
if not is_admin:
query = query.filter((Pipeline.owner_id == current_user_id) | (Pipeline.owner_id.is_(None)))
elif owner_id:
query = query.filter(Pipeline.owner_id == owner_id)
elif owner_id:
query = query.filter(Pipeline.owner_id == owner_id)
rows = query.order_by(Pipeline.id).offset(offset).limit(limit).all()
return [_pipeline_from_record(p, db) for p in rows]
@strawberry.field(description="Fetch a single pipeline by ID.")
def pipeline(self, info: strawberry.types.Info, id: int) -> PipelineType | None:
"""Return one pipeline by its primary key, or *null* if not found."""
db, user = _get_db_and_user(info)
_require_auth(user)
row = db.query(Pipeline).filter(Pipeline.id == id).first()
if row is None:
return None
if settings.auth_enabled and user:
is_admin = user.get("is_admin", False)
current_user_id = _get_current_user_id(user)
if not is_admin and row.owner_id is not None and row.owner_id != current_user_id:
return None
return _pipeline_from_record(row, db)
@strawberry.field(description="List non-sensitive application settings (admin only).")
def settings(
self,
info: strawberry.types.Info,
limit: int = 50,
offset: int = 0,
) -> list[SettingType]:
"""Return application settings stored in the database.
Sensitive keys (API secrets, passwords, etc.) are automatically
excluded. Requires admin privileges when auth is enabled.
"""
db, user = _get_db_and_user(info)
_require_admin(user)
limit = max(1, min(limit, 200))
offset = max(0, offset)
rows = (
db.query(ApplicationSettings)
.filter(ApplicationSettings.key.notin_(_SENSITIVE_SETTING_KEYS))
.order_by(ApplicationSettings.key)
.offset(offset)
.limit(limit)
.all()
)
return [_setting_from_record(r) for r in rows]
@strawberry.field(description="List user profiles (admin only).")
def users(
self,
info: strawberry.types.Info,
limit: int = 20,
offset: int = 0,
) -> list[UserType]:
"""Return a paginated list of user profiles. Requires admin privileges."""
db, user = _get_db_and_user(info)
_require_admin(user)
limit = max(1, min(limit, 100))
offset = max(0, offset)
rows = db.query(UserProfile).order_by(UserProfile.user_id).offset(offset).limit(limit).all()
return [_user_from_profile(r) for r in rows]
@strawberry.field(description="Fetch a user profile by user_id (admin only).")
def user(self, info: strawberry.types.Info, user_id: str) -> UserType | None:
"""Return one user profile by *user_id*, or *null* if not found."""
db, user = _get_db_and_user(info)
_require_admin(user)
row = db.query(UserProfile).filter(UserProfile.user_id == user_id).first()
return _user_from_profile(row) if row else None
# ---------------------------------------------------------------------------
# Schema and router
# ---------------------------------------------------------------------------
schema = strawberry.Schema(query=Query)
async def get_graphql_context(
request: Request,
db: Annotated[Session, Depends(get_db)],
) -> dict[str, Any]:
"""Build the per-request context injected into every resolver."""
try:
user = get_current_user(request)
except Exception:
logger.debug("Could not resolve current user for GraphQL context", exc_info=True)
user = None
return {"request": request, "db": db, "user": user}
graphql_router = GraphQLRouter(
schema,
context_getter=get_graphql_context,
graphql_ide="graphiql",
)
+2
View File
@@ -16,6 +16,7 @@ from starlette.middleware.trustedhost import TrustedHostMiddleware
from uvicorn.middleware.proxy_headers import ProxyHeadersMiddleware
from app.api import router as api_router
from app.api.graphql_api import graphql_router
from app.api.local_auth import router as local_auth_router
from app.auth import router as auth_router
from app.config import settings
@@ -298,3 +299,4 @@ app.include_router(files_router) # Explicitly include the files router
app.include_router(auth_router)
app.include_router(local_auth_router)
app.include_router(api_router, prefix="/api")
app.include_router(graphql_router, prefix="/graphql")
+119
View File
@@ -2132,3 +2132,122 @@ Return basic profile information for the authenticated user.
"is_admin": false
}
```
---
## GraphQL API
DocuElevate exposes a GraphQL API at `/graphql` alongside the REST API. It
supports flexible queries with field selection, making it ideal for dashboards
and integrations that only need a subset of the available data.
### Endpoint
| Method | URL | Description |
|--------|-----|-------------|
| `POST` | `/graphql` | Execute a GraphQL query or mutation |
| `GET` | `/graphql` | Open the GraphiQL interactive playground |
### Authentication
The GraphQL endpoint honours the same authentication rules as the REST API:
- **`AUTH_ENABLED=False`** (default, single-user mode): all queries are
allowed without credentials.
- **`AUTH_ENABLED=True`** (multi-user mode): a valid session cookie **or**
an `Authorization: Bearer <token>` API token is required. Admin-only
queries (settings, users) additionally require the `is_admin` flag.
### Available Queries
| Field | Returns | Notes |
|-------|---------|-------|
| `documents(ownerId, limit, offset)` | `[DocumentType]` | Paginated list of documents |
| `document(id)` | `DocumentType` | Single document by primary key |
| `pipelines(ownerId, limit, offset)` | `[PipelineType]` | Paginated list of pipelines with steps |
| `pipeline(id)` | `PipelineType` | Single pipeline by primary key |
| `settings(limit, offset)` | `[SettingType]` | Non-sensitive app settings (**admin only**) |
| `users(limit, offset)` | `[UserType]` | User profiles (**admin only**) |
| `user(userId)` | `UserType` | Single user profile (**admin only**) |
> **Note:** Sensitive configuration keys (API secrets, passwords, tokens) are
> automatically excluded from the `settings` query regardless of the caller's
> privilege level.
### GraphiQL Playground
Navigate to `http://<your-instance>/graphql` in a browser to open the
interactive GraphiQL IDE, which provides schema documentation, auto-complete,
and the ability to run queries directly.
### Example Queries
**List recent documents:**
```graphql
{
documents(limit: 5) {
id
originalFilename
mimeType
fileSize
documentTitle
createdAt
}
}
```
**Fetch a pipeline with its steps:**
```graphql
{
pipeline(id: 1) {
id
name
description
isDefault
isActive
steps {
position
stepType
label
enabled
}
}
}
```
**List application settings (admin only):**
```graphql
{
settings {
key
value
updatedAt
}
}
```
**List user profiles (admin only):**
```graphql
{
users(limit: 10) {
userId
displayName
subscriptionTier
isBlocked
}
}
```
**Using variables:**
```graphql
query GetDocument($id: Int!) {
document(id: $id) {
id
originalFilename
documentTitle
isDuplicate
ocrQualityScore
}
}
```
Variables: `{ "id": 42 }`
+4 -1
View File
@@ -51,4 +51,7 @@ meilisearch>=0.31.0 # Full-text search engine client
stripe>=7.0.0,<15.0.0 # Stripe billing SDK (MIT license)
# Error and performance monitoring
sentry-sdk[fastapi,celery,sqlalchemy]>=2.20.0,<3.0.0
sentry-sdk[fastapi,celery,sqlalchemy]>=2.20.0,<3.0.0
# GraphQL API
strawberry-graphql[fastapi]>=0.243.0,<1.0.0
+419
View File
@@ -0,0 +1,419 @@
"""
Tests for the GraphQL API endpoint at /graphql.
Covers:
- Schema introspection (endpoint availability + GraphiQL)
- Query: documents (list, single, auth-gated)
- Query: pipelines (list, single)
- Query: settings (admin-only)
- Query: users (admin-only)
- Pagination and limit clamping
- Sensitive setting keys are excluded
"""
from __future__ import annotations
from unittest.mock import patch
import pytest
from fastapi.testclient import TestClient
from app.models import ApplicationSettings, FileRecord, Pipeline, PipelineStep, UserProfile
# ---------------------------------------------------------------------------
# Helpers
# ---------------------------------------------------------------------------
def gql(client: TestClient, query: str, variables: dict | None = None) -> dict:
"""Execute a GraphQL POST request and return the parsed JSON body."""
payload: dict = {"query": query}
if variables:
payload["variables"] = variables
response = client.post("/graphql", json=payload)
assert response.status_code == 200, f"Unexpected status {response.status_code}: {response.text}"
return response.json()
# ---------------------------------------------------------------------------
# Fixtures
# ---------------------------------------------------------------------------
@pytest.fixture()
def file_record(db_session) -> FileRecord:
rec = FileRecord(
owner_id="user1",
original_filename="invoice.pdf",
local_filename="/workdir/tmp/invoice.pdf",
file_size=1024,
mime_type="application/pdf",
filehash="abc123",
)
db_session.add(rec)
db_session.commit()
db_session.refresh(rec)
return rec
@pytest.fixture()
def pipeline_record(db_session) -> Pipeline:
p = Pipeline(
owner_id="user1",
name="Test Pipeline",
description="A pipeline for tests",
is_default=False,
is_active=True,
)
db_session.add(p)
db_session.commit()
db_session.refresh(p)
step = PipelineStep(
pipeline_id=p.id,
position=0,
step_type="ocr",
label="Run OCR",
enabled=True,
)
db_session.add(step)
db_session.commit()
return p
@pytest.fixture()
def setting_record(db_session) -> ApplicationSettings:
s = ApplicationSettings(key="max_upload_size", value="104857600")
db_session.add(s)
db_session.commit()
db_session.refresh(s)
return s
@pytest.fixture()
def user_profile(db_session) -> UserProfile:
profile = UserProfile(
user_id="user1",
display_name="Test User",
is_blocked=False,
subscription_tier="free",
onboarding_completed=False,
)
db_session.add(profile)
db_session.commit()
db_session.refresh(profile)
return profile
# ---------------------------------------------------------------------------
# Tests: endpoint availability
# ---------------------------------------------------------------------------
@pytest.mark.unit
class TestGraphQLEndpoint:
"""Verify the /graphql endpoint is reachable and introspectable."""
def test_graphql_post_exists(self, client: TestClient):
"""POST /graphql returns 200 for a valid introspection query."""
result = gql(client, "{ __schema { queryType { name } } }")
assert "data" in result
assert result["data"]["__schema"]["queryType"]["name"] == "Query"
def test_graphql_get_returns_graphiql(self, client: TestClient):
"""GET /graphql returns the GraphiQL playground HTML."""
response = client.get("/graphql", headers={"Accept": "text/html"})
assert response.status_code == 200
assert "graphiql" in response.text.lower() or "graphql" in response.text.lower()
def test_graphql_schema_has_expected_types(self, client: TestClient):
"""Schema exposes DocumentType, PipelineType, SettingType, UserType."""
result = gql(
client,
"""
{
__schema {
types { name }
}
}
""",
)
type_names = {t["name"] for t in result["data"]["__schema"]["types"]}
for expected in ("DocumentType", "PipelineType", "SettingType", "UserType"):
assert expected in type_names, f"{expected} not found in schema"
def test_graphql_query_fields(self, client: TestClient):
"""Root Query has documents, document, pipelines, pipeline, settings, users, user fields."""
result = gql(
client,
"""
{
__type(name: "Query") {
fields { name }
}
}
""",
)
field_names = {f["name"] for f in result["data"]["__type"]["fields"]}
for expected in ("documents", "document", "pipelines", "pipeline", "settings", "users", "user"):
assert expected in field_names, f"Query field '{expected}' missing from schema"
# ---------------------------------------------------------------------------
# Tests: documents queries
# ---------------------------------------------------------------------------
@pytest.mark.integration
class TestDocumentsQuery:
"""Tests for the documents and document queries."""
def test_list_documents_empty(self, client: TestClient):
result = gql(client, "{ documents { id originalFilename } }")
assert "errors" not in result
assert result["data"]["documents"] == []
def test_list_documents_returns_records(self, client: TestClient, file_record: FileRecord):
result = gql(client, "{ documents { id originalFilename mimeType fileSize } }")
assert "errors" not in result
docs = result["data"]["documents"]
assert len(docs) == 1
assert docs[0]["id"] == file_record.id
assert docs[0]["originalFilename"] == "invoice.pdf"
assert docs[0]["mimeType"] == "application/pdf"
assert docs[0]["fileSize"] == 1024
def test_get_single_document(self, client: TestClient, file_record: FileRecord):
result = gql(
client,
"query($id: Int!) { document(id: $id) { id originalFilename } }",
variables={"id": file_record.id},
)
assert "errors" not in result
assert result["data"]["document"]["id"] == file_record.id
def test_get_nonexistent_document_returns_null(self, client: TestClient):
result = gql(client, "{ document(id: 99999) { id } }")
assert "errors" not in result
assert result["data"]["document"] is None
def test_documents_pagination(self, client: TestClient, db_session):
for i in range(5):
db_session.add(
FileRecord(
owner_id="user1",
original_filename=f"doc{i}.pdf",
local_filename=f"/workdir/tmp/doc{i}.pdf",
file_size=100,
filehash=f"hash{i}",
)
)
db_session.commit()
result_page1 = gql(client, "{ documents(limit: 2, offset: 0) { id } }")
result_page2 = gql(client, "{ documents(limit: 2, offset: 2) { id } }")
assert "errors" not in result_page1
assert "errors" not in result_page2
assert len(result_page1["data"]["documents"]) == 2
assert len(result_page2["data"]["documents"]) == 2
def test_documents_limit_clamped_to_100(self, client: TestClient, db_session):
# Requesting more than 100 should be silently clamped to 100
for i in range(5):
db_session.add(
FileRecord(
owner_id="user1",
original_filename=f"big{i}.pdf",
local_filename=f"/workdir/tmp/big{i}.pdf",
file_size=100,
filehash=f"bighash{i}",
)
)
db_session.commit()
result = gql(client, "{ documents(limit: 999) { id } }")
assert "errors" not in result
# Just verify it doesn't error and returns something
assert isinstance(result["data"]["documents"], list)
def test_documents_filter_by_owner(self, client: TestClient, db_session):
db_session.add(
FileRecord(
owner_id="alice",
original_filename="alice.pdf",
local_filename="/workdir/tmp/alice.pdf",
file_size=100,
filehash="alicehash",
)
)
db_session.add(
FileRecord(
owner_id="bob",
original_filename="bob.pdf",
local_filename="/workdir/tmp/bob.pdf",
file_size=200,
filehash="bobhash",
)
)
db_session.commit()
result = gql(client, '{ documents(ownerId: "alice") { id originalFilename } }')
assert "errors" not in result
docs = result["data"]["documents"]
assert all(d["originalFilename"] == "alice.pdf" for d in docs)
# ---------------------------------------------------------------------------
# Tests: pipelines queries
# ---------------------------------------------------------------------------
@pytest.mark.integration
class TestPipelinesQuery:
"""Tests for the pipelines and pipeline queries."""
def test_list_pipelines_empty(self, client: TestClient):
result = gql(client, "{ pipelines { id name } }")
assert "errors" not in result
assert result["data"]["pipelines"] == []
def test_list_pipelines_with_steps(self, client: TestClient, pipeline_record: Pipeline):
result = gql(
client,
"""
{
pipelines {
id name description isDefault isActive
steps { id stepType position enabled }
}
}
""",
)
assert "errors" not in result
pipelines = result["data"]["pipelines"]
assert len(pipelines) == 1
assert pipelines[0]["name"] == "Test Pipeline"
assert len(pipelines[0]["steps"]) == 1
assert pipelines[0]["steps"][0]["stepType"] == "ocr"
def test_get_single_pipeline(self, client: TestClient, pipeline_record: Pipeline):
result = gql(
client,
"query($id: Int!) { pipeline(id: $id) { id name steps { stepType } } }",
variables={"id": pipeline_record.id},
)
assert "errors" not in result
assert result["data"]["pipeline"]["id"] == pipeline_record.id
assert result["data"]["pipeline"]["steps"][0]["stepType"] == "ocr"
def test_get_nonexistent_pipeline_returns_null(self, client: TestClient):
result = gql(client, "{ pipeline(id: 99999) { id } }")
assert "errors" not in result
assert result["data"]["pipeline"] is None
# ---------------------------------------------------------------------------
# Tests: settings query (admin-only)
# ---------------------------------------------------------------------------
@pytest.mark.unit
class TestSettingsQuery:
"""Tests for the settings query."""
def test_settings_returns_data_when_no_auth(self, client: TestClient, setting_record: ApplicationSettings):
"""When AUTH_ENABLED=False, settings are accessible (no auth required)."""
result = gql(client, "{ settings { key value } }")
assert "errors" not in result
keys = [s["key"] for s in result["data"]["settings"]]
assert "max_upload_size" in keys
def test_sensitive_settings_excluded(self, client: TestClient, db_session):
"""Sensitive setting keys must never appear in the response."""
sensitive_keys = [
"openai_api_key",
"session_secret",
"azure_ai_key",
"smtp_password",
]
for key in sensitive_keys:
db_session.add(ApplicationSettings(key=key, value="super-secret"))
db_session.commit()
result = gql(client, "{ settings { key value } }")
assert "errors" not in result
returned_keys = {s["key"] for s in result["data"]["settings"]}
for key in sensitive_keys:
assert key not in returned_keys, f"Sensitive key '{key}' was returned by GraphQL settings query"
def test_settings_auth_required_when_auth_enabled(self, client: TestClient):
"""When AUTH_ENABLED=True and no user, settings query must return an error."""
from app.config import settings as app_settings
with patch.object(app_settings, "auth_enabled", True):
result = gql(client, "{ settings { key } }")
# Should have errors because no user is authenticated
assert "errors" in result
# ---------------------------------------------------------------------------
# Tests: users query (admin-only)
# ---------------------------------------------------------------------------
@pytest.mark.unit
class TestUsersQuery:
"""Tests for the users and user queries."""
def test_users_returns_profiles_when_no_auth(self, client: TestClient, user_profile: UserProfile):
"""When AUTH_ENABLED=False, users are accessible."""
result = gql(client, "{ users { userId displayName subscriptionTier } }")
assert "errors" not in result
users = result["data"]["users"]
assert any(u["userId"] == "user1" for u in users)
def test_get_user_by_id(self, client: TestClient, user_profile: UserProfile):
result = gql(
client,
'query { user(userId: "user1") { userId displayName isBlocked } }',
)
assert "errors" not in result
assert result["data"]["user"]["userId"] == "user1"
assert result["data"]["user"]["displayName"] == "Test User"
assert result["data"]["user"]["isBlocked"] is False
def test_get_nonexistent_user_returns_null(self, client: TestClient):
result = gql(client, '{ user(userId: "nobody") { userId } }')
assert "errors" not in result
assert result["data"]["user"] is None
def test_users_auth_required_when_auth_enabled(self, client: TestClient):
"""When AUTH_ENABLED=True and no user, users query must return an error."""
from app.config import settings as app_settings
with patch.object(app_settings, "auth_enabled", True):
result = gql(client, "{ users { userId } }")
assert "errors" in result
# ---------------------------------------------------------------------------
# Tests: auth enforcement
# ---------------------------------------------------------------------------
@pytest.mark.unit
class TestGraphQLAuth:
"""Verify auth is enforced for all queries when AUTH_ENABLED=True."""
def test_documents_auth_required_when_auth_enabled(self, client: TestClient):
from app.config import settings as app_settings
with patch.object(app_settings, "auth_enabled", True):
result = gql(client, "{ documents { id } }")
assert "errors" in result
def test_pipelines_auth_required_when_auth_enabled(self, client: TestClient):
from app.config import settings as app_settings
with patch.object(app_settings, "auth_enabled", True):
result = gql(client, "{ pipelines { id } }")
assert "errors" in result