diff --git a/.env.example b/.env.example index 2c1ab31c..dae53ee5 100644 --- a/.env.example +++ b/.env.example @@ -13,6 +13,8 @@ # ============================================================================= # --- Django ----------------------------------------------------------------- +# REQUIRED: startup refuses to boot while this is empty or still `change-me`. +# Generate one with: openssl rand -hex 32 DJANGO_SECRET_KEY=change-me DJANGO_DEBUG=false DJANGO_ALLOWED_HOSTS=* @@ -88,6 +90,31 @@ WORKER_POOL=cpu # skips if demo runs already exist. Set to false to opt out. SEED_DEMO_AUDITS=true +# --- Chat (Open WebUI, optional) ------------------------------------------------ +# Embeds Open WebUI at /chat/, signed in as the Studio user (workspace admins +# become Open WebUI admins). Off by default — uncomment both lines to enable: +# the first makes the web container serve /chat/, the second makes +# `docker compose up -d` start the chat containers without an extra --profile +# flag. See docs/chat.md. +#SIMPLEAUDIT_CHAT=docker +#COMPOSE_PROFILES=chat +# The origin the browser opens (the chat proxy). Use your own hostname behind +# TLS; on a separate subdomain also set SESSION_COOKIE_DOMAIN. +SIMPLEAUDIT_CHAT_URL=http://localhost:8801 +# Where signed-out users are sent back to. +SIMPLEAUDIT_STUDIO_URL=http://localhost:8000 +# Optional: export Open WebUI's spans to Studio's OTLP listener. Off by default. +# Exports unauthenticated (matches an enabled "none" credential). For a +# basic/bearer credential, set OTEL_BASIC_AUTH_* to match it. See docs/chat.md. +#SIMPLEAUDIT_CHAT_OTLP=true +#SIMPLEAUDIT_CHAT_OTLP_ENDPOINT=http://web:8000 +#OTEL_BASIC_AUTH_USERNAME= +#OTEL_BASIC_AUTH_PASSWORD= +# Optional: switch Studio's OTLP listener (POST /otlp/v1/traces + credential +# management) on or off. ON by default (preserves existing behavior). Set to +# "off" to 404 the /otlp/* and /api/otlp/* routes and hide the OTLP button. +#SIMPLEAUDIT_OTLP=off + # --- Optional observability ----------------------------------------------------- # --- Sentry error tracking & tracing (leave empty to disable) ------------------- diff --git a/.github/agents/orchestrator.agent.md b/.github/agents/orchestrator.agent.md new file mode 100644 index 00000000..d8d5b48b --- /dev/null +++ b/.github/agents/orchestrator.agent.md @@ -0,0 +1,47 @@ +--- +name: Orchestrator +description: Parallel-first coding orchestrator +tools: ['agent', 'edit', 'read', 'search', 'execute'] +agents: ['Researcher', 'Verifier', 'Test Reviewer', 'Regression Reviewer'] +--- + +You are the main engineering orchestrator. + +Default behavior: +- Decompose non-trivial tasks into independent workstreams. +- Run independent research/review tasks in parallel whenever possible. +- Prefer parallel subagents over doing sequential investigation yourself. +- Keep the main context focused on decisions and integration. +- Do not delegate trivial tasks where coordination overhead exceeds the work. + +For implementation: +1. First identify independent components. +2. Launch parallel subagents for codebase research, dependency analysis, + test discovery, and alternative implementation approaches. +3. Integrate the best findings yourself. +4. After editing, launch independent reviewers in parallel: + - correctness + - test coverage + - regression risk + - hallucinated APIs / assumptions +5. Run tests and inspect actual outputs. +6. Fix issues found by reviewers. +7. Repeat verification until there are no actionable failures. + +Never trust another agent's factual claim about the repository unless it +provides file references or you verify it yourself. + +Prefer evidence from: +- repository contents +- compiler/type checker +- test output +- runtime output +- official documentation + +If uncertain, investigate rather than guessing. + +Keep working autonomously until: +- the requested result is implemented, +- tests/checks have been run, +- significant reviewer findings have been addressed, +- or a genuine blocker requires user input. diff --git a/.github/agents/regression-reviewer.agent.md b/.github/agents/regression-reviewer.agent.md new file mode 100644 index 00000000..0c8bab87 --- /dev/null +++ b/.github/agents/regression-reviewer.agent.md @@ -0,0 +1,21 @@ +--- +name: Regression Reviewer +description: Assess regression risk of proposed changes +user-invocable: false +tools: ['read', 'search', 'execute'] +--- + +Assess the regression risk of the proposed changes. + +Check: +- existing behavior that depends on the changed code +- callers and consumers that may break +- shared state, configuration, or schema changes +- edge cases the change may have disturbed + +Use repository evidence and executable checks whenever possible. + +Return only: +1. confirmed regression risks +2. evidence +3. concrete mitigations diff --git a/.github/agents/researcher.agent.md b/.github/agents/researcher.agent.md new file mode 100644 index 00000000..77c352a2 --- /dev/null +++ b/.github/agents/researcher.agent.md @@ -0,0 +1,13 @@ +--- +name: Researcher +description: Fast codebase reconnaissance +user-invocable: false +tools: ['read', 'search'] +--- + +Investigate the assigned question deeply but do not edit files. + +Search broadly, identify relevant files and existing patterns, +and return concise findings with exact file references. + +Do not speculate when repository evidence is available. diff --git a/.github/agents/test-reviewer.agent.md b/.github/agents/test-reviewer.agent.md new file mode 100644 index 00000000..1d06f6b9 --- /dev/null +++ b/.github/agents/test-reviewer.agent.md @@ -0,0 +1,21 @@ +--- +name: Test Reviewer +description: Review test coverage for changed behavior +user-invocable: false +tools: ['read', 'search', 'execute'] +--- + +Review the tests for the changed behavior. + +Check: +- tests exist for the new or changed behavior +- tests assert the right outcomes, not just that code runs +- edge cases and failure paths are covered +- the test suite actually passes when run + +Run the relevant tests and report actual output. + +Return only: +1. confirmed gaps or failures +2. evidence +3. concrete fixes diff --git a/.github/agents/verifier.agent.md b/.github/agents/verifier.agent.md new file mode 100644 index 00000000..d0ce0686 --- /dev/null +++ b/.github/agents/verifier.agent.md @@ -0,0 +1,25 @@ +--- +name: Verifier +description: Verify implementation claims and detect hallucinations +user-invocable: false +tools: ['read', 'search', 'execute'] +--- + +Independently verify the proposed solution. + +Do not assume another agent's claims are correct. + +Check: +- referenced files and symbols actually exist +- APIs and function signatures are real +- dependencies actually expose the claimed features +- tests exercise the changed behavior +- implementation matches the original request +- no placeholder or speculative code remains + +Use repository evidence and executable checks whenever possible. + +Return only: +1. confirmed problems +2. evidence +3. concrete fixes diff --git a/.github/prompts/parallel-orchestrate.prompt.md b/.github/prompts/parallel-orchestrate.prompt.md new file mode 100644 index 00000000..dfc10170 --- /dev/null +++ b/.github/prompts/parallel-orchestrate.prompt.md @@ -0,0 +1,21 @@ +--- +name: parallel-orchestrate +description: Parallel-first orchestration convention for substantial tasks +--- + +Use parallel subagents aggressively for independent work. Optimize for +wall-clock time, not token count. Treat model calls as cheap. Keep yourself +as the coordinator and integrator. Verify claims with repository evidence +and executable checks before concluding. + +For every substantial task: +1. Launch 2-3 Researcher subagents in parallel for independent angles + (simplest implementation, architecture-compatible implementation, + hidden risks). +2. Integrate the best findings and implement. +3. Launch Verifier, Test Reviewer, and Regression Reviewer in parallel. +4. Fix confirmed findings and re-verify until there are no actionable + failures. + +Keep the delegation tree wide, not deep: orchestrator -> workers/reviewers. +Do not nest subagents more than one level. diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index b8ac1de2..fecb3366 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -28,17 +28,28 @@ jobs: with: python-version: "3.12" - uses: astral-sh/setup-uv@v7 + - name: Clone core (local-dev path dependency) + # pyproject.toml points simpleaudit at ../SimpleAudit (editable) so the + # studio can use the tracing layer not yet in the published wheel. + run: git clone --depth 1 https://github.com/kelkalot/simpleaudit.git ../SimpleAudit - name: Install dependencies run: uv sync --frozen --extra dev - name: Lint (ruff) run: uv run ruff check . - - name: Run tests (parallel) - run: uv run python manage.py test infra --parallel auto --exclude-tag embedded_hatchet + - name: Cache testmon dependency data + uses: actions/cache@v4 + with: + path: .testmondata + key: testmon-${{ github.sha }} + restore-keys: | + testmon- + - name: Run tests (parallel, affected-only via testmon) + run: uv run pytest --testmon -n auto -m "not embedded_hatchet" - name: Check if embedded Hatchet files changed id: hatchet-changed run: | CHANGED=$(git diff --name-only origin/main...HEAD -- infra/minimal_config.py infra/worker.py 2>/dev/null || true) - if [ -n "$CHANGED" ]; then echo "changed=true" >> "$GITHUB_OUTPUT"; else echo "changed=false" >> "$GITHUB_OUTPUT"; fi + if [ -n "$CHANGED"]; then echo "changed=true" >> "$GITHUB_OUTPUT"; else echo "changed=false" >> "$GITHUB_OUTPUT"; fi - name: Run tests (embedded Hatchet, serial) if: steps.hatchet-changed.outputs.changed == 'true' - run: uv run python manage.py test infra.tests.test_minimal_config.TestEmbeddedHatchetLifecycle + run: uv run pytest -m "embedded_hatchet" diff --git a/.github/workflows/docker.yml b/.github/workflows/docker.yml index 27805b8c..b15c034d 100644 --- a/.github/workflows/docker.yml +++ b/.github/workflows/docker.yml @@ -38,6 +38,10 @@ jobs: runs-on: ubuntu-latest steps: - uses: actions/checkout@v4 + - name: Clone core (local-dev path dependency) + # pyproject.toml points simpleaudit at ../SimpleAudit (editable) so the + # studio can use the tracing layer not yet in the published wheel. + run: git clone --depth 1 https://github.com/kelkalot/simpleaudit.git ../SimpleAudit - name: Build HF Space / demo image run: docker build -t simpleaudit-studio:ci . - name: Build Compose image diff --git a/.gitignore b/.gitignore index e2cf8275..971cdb16 100644 --- a/.gitignore +++ b/.gitignore @@ -44,3 +44,4 @@ app.pid # Data artifacts (prototype) data/ +.testmondata* diff --git a/Dockerfile b/Dockerfile index 9f3ec9ad..db9133d4 100644 --- a/Dockerfile +++ b/Dockerfile @@ -25,13 +25,18 @@ WORKDIR /app RUN apt-get update \ && apt-get install -y --no-install-recommends \ curl \ + git \ && rm -rf /var/lib/apt/lists/* # --- Python dependencies ----------------------------------------------------- # pyproject.toml is the single source of truth; uv.lock pins exact versions. # uv sync creates /app/.venv; the PATH update keeps the `python` entrypoint. +# The core (simpleaudit) is a local path dependency (../SimpleAudit) so the +# studio can use the tracing layer not yet in the published wheel. Clone it +# into the build context before uv sync. COPY pyproject.toml uv.lock README.md ./ -RUN pip install uv \ +RUN git clone --depth 1 https://github.com/kelkalot/simpleaudit.git /SimpleAudit \ + && pip install uv \ && uv sync --frozen --no-install-project --no-dev ENV PATH="/app/.venv/bin:$PATH" diff --git a/README.md b/README.md index f89f7a49..9dd03239 100644 --- a/README.md +++ b/README.md @@ -77,11 +77,43 @@ Every start applies migrations and makes sure the admin (a superuser) and defaul ### Tests and lint +Tests run under **pytest** (via `pytest-django`). Your existing +`django.test.TestCase` classes run unchanged. Tests are layered: run the +**fast** set while coding (skips the slow integration modules tagged `slow`), +and the **full** set before you commit or open a PR. + ```bash -SIMPLEAUDIT_LOCAL_SQLITE=1 uv run manage.py test infra --exclude-tag embedded_hatchet +# Fast — unit + light integration, for the dev loop (~50s, skips the slow modules) +uv run pytest -n auto -m "not slow and not embedded_hatchet" + +# Full — everything, for before commit / PR +uv run pytest -n auto -m "not embedded_hatchet" + +# Affected-only — run just the tests touched by your changed code (needs a +# prior run to build .testmondata; CI caches it) +uv run pytest --testmon -n auto -m "not slow and not embedded_hatchet" + +# Lint uv run ruff check . + +# Target a single module, class, or test to iterate faster +uv run pytest infra/tests/test_workspaces.py # one module +uv run pytest infra/tests/test_workspaces.py::WorkspaceTests # one class +uv run pytest infra/tests/test_workspaces.py::WorkspaceTests::test_create # one test ``` +**Tagging** + +- `slow` — heavy integration modules (experiments, monitors, judges, engine + integration, full API lifecycle). Skipped by the fast command, always run in + CI. Add `@tag("slow")` to a class to move a slow test out of the fast loop. +- `embedded_hatchet` — starts a real embedded Hatchet worker; runs serially in + CI only when the relevant files change. + +The Django `@tag("...")` values are mirrored onto pytest markers by +`conftest.py`, so `-m "not slow"` works the same as +`manage.py test --exclude-tag slow`. + ### Other setups ```bash @@ -99,14 +131,35 @@ For teams or multi-user setups, use Docker Compose: git clone https://github.com/SushantGautam/SimpleAuditStudio cd SimpleAuditStudio cp .env.example .env -# edit POSTGRES_PASSWORD and BOOTSTRAP_PASSWORD at minimum +# edit DJANGO_SECRET_KEY, POSTGRES_PASSWORD and BOOTSTRAP_PASSWORD at minimum — +# startup refuses to boot while any of them is empty or still `change-me` docker compose up -d ``` -Services: Web UI (:8000), PostgreSQL, Hatchet queue (:8888), Worker. Optional profile: `--profile mock` (mock model API). +Services: Web UI (:8000), PostgreSQL, Hatchet queue (:8888), Worker. Chat (Open WebUI) is included via `.env`; optional profile `--profile mock` adds a mock model API. See [docs/deployment.md](docs/deployment.md) for production hardening, backups, and upgrades. +## 💬 Chat + +SimpleAudit Studio embeds [Open WebUI](https://openwebui.com) at `/chat/`, signed +in as your Studio user — workspace admins become Open WebUI admins. + +Chat is opt-in. The local one-liner bundles it by default (pass +`--disable-chat` to turn it off); Docker Compose leaves it out unless `.env` +says otherwise — uncomment `SIMPLEAUDIT_CHAT` and `COMPOSE_PROFILES` in +`.env.example` to include it: + +```bash +uvx simpleaudit-studio # chat included +uvx simpleaudit-studio --disable-chat # without it + +docker compose up -d # chat only if enabled in .env +``` + +See [docs/chat.md](docs/chat.md) for how single sign-on works and what must stay +private. + ## ✨ What You Can Do - Build versioned scenario sets and register OpenAI-compatible models diff --git a/accounts/serializers.py b/accounts/serializers.py index 398905bc..503e9f55 100644 --- a/accounts/serializers.py +++ b/accounts/serializers.py @@ -33,7 +33,11 @@ def validate_password(self, value): return value def create(self, validated_data): - return User.objects.create_user(**validated_data) + user = User.objects.create_user(**validated_data) + from accounts.services import grant_default_project + + grant_default_project(user) + return user class WorkspaceItemSerializer(serializers.ModelSerializer): diff --git a/accounts/services.py b/accounts/services.py index 03e33dc9..3579e049 100644 --- a/accounts/services.py +++ b/accounts/services.py @@ -79,6 +79,23 @@ def bootstrap_admin_and_default_project( DEFAULT_PROJECT_SLUG = "default" +def grant_default_project(user) -> Project | None: + """Give a new user a viewer membership in the shared 'default' workspace. + + Every user-creation path (admin add-user, self-registration, WorkOS magic + auth, demo signup) calls this so a fresh account always has a project to + land in — otherwise ``request.project`` resolves to ``None`` and the UI + 500s. Idempotent: an existing membership is left untouched. Returns the + default project, or ``None`` when it does not exist. + """ + project = Project.objects.filter(slug=DEFAULT_PROJECT_SLUG).first() + if project: + ProjectMembership.objects.get_or_create( + project=project, user=user, defaults={"role": ProjectMembership.Role.VIEWER} + ) + return project + + def ensure_project_access(user, project) -> bool: """Return True if the user may view this project's content. @@ -263,13 +280,15 @@ def admin_create_user(*, admin_user, username: str, email: str = "", password: s from django.contrib.auth.password_validation import validate_password validate_password(password) - return User.objects.create_user( + user = User.objects.create_user( username=clean_username, email=clean_email, password=password, first_name=(first_name or "").strip(), last_name=(last_name or "").strip(), ) + grant_default_project(user) + return user @transaction.atomic diff --git a/audits/migrations/0005_auditrun_trace_config.py b/audits/migrations/0005_auditrun_trace_config.py new file mode 100644 index 00000000..c1047218 --- /dev/null +++ b/audits/migrations/0005_auditrun_trace_config.py @@ -0,0 +1,17 @@ +# Generated by Django 5.2.4 on 2026-10-01 19:05 + +from django.db import migrations, models + + +class Migration(migrations.Migration): + dependencies = [ + ("audits", "0004_judge_model_per_run"), + ] + + operations = [ + migrations.AddField( + model_name="auditrun", + name="trace_config", + field=models.JSONField(blank=True, default=dict), + ), + ] diff --git a/audits/models.py b/audits/models.py index 35902d42..4f141518 100644 --- a/audits/models.py +++ b/audits/models.py @@ -38,6 +38,10 @@ class Status(models.TextChoices): auditor_config_snapshot = models.JSONField() judge_config_snapshot = models.JSONField() generation_parameters_snapshot = models.JSONField() + # Optional trace acquisition config (Promptfoo parity): {"mode": "builtin"|"tempo", ...}. + # Empty/absent = no tracing (the normal black-box path). Frozen at run creation + # like the other snapshots; the worker hands it to the engine's tracing layer. + trace_config = models.JSONField(default=dict, blank=True) simpleaudit_version = models.CharField(max_length=120) git_commit = models.CharField(max_length=120) runtime_metadata = models.JSONField(default=dict, blank=True) diff --git a/audits/monitors.py b/audits/monitors.py index 669f3fce..52e165b9 100644 --- a/audits/monitors.py +++ b/audits/monitors.py @@ -10,7 +10,6 @@ from __future__ import annotations import logging -import math from datetime import UTC, timedelta from django.db import transaction @@ -348,6 +347,27 @@ def create_monitor(*, project, user, name: str, run: dict, repeat: dict, experim return monitor +def due_monitors(now): + """Claim every enabled monitor that is due, locking only the monitor rows. + + ``of=("self",)`` is required, not a refinement. ``last_run`` and ``created_by`` + are both nullable, so ``select_related`` joins them with a LEFT OUTER JOIN, and + PostgreSQL rejects ``FOR UPDATE`` against the nullable side of an outer join: + + FOR UPDATE cannot be applied to the nullable side of an outer join + + Without ``of``, every tick raises that on PostgreSQL -- so no monitor ever runs + on a production database, while the sweeper logs the failure and carries on. + SQLite omits ``FOR UPDATE`` entirely, which is why the test suite stayed green. + """ + return ( + Monitor.objects.select_for_update(of=("self",), skip_locked=True) + .filter(enabled=True, next_run_at__lte=now) + .select_related("last_run", "project", "created_by") + .order_by("next_run_at") + ) + + def run_due_monitors(now=None) -> list[int]: """Launch every enabled monitor whose ``next_run_at`` has passed. @@ -362,12 +382,7 @@ def run_due_monitors(now=None) -> list[int]: now = now or timezone.now() created: list[AuditRun] = [] with transaction.atomic(): - due = ( - Monitor.objects.select_for_update(skip_locked=True) - .filter(enabled=True, next_run_at__lte=now) - .select_related("last_run", "project", "created_by") - .order_by("next_run_at") - ) + due = due_monitors(now) for monitor in due: monitor.next_run_at = next_after(monitor, now) monitor.last_tick_at = now @@ -449,25 +464,17 @@ def pass_counts(run_ids: list[int]) -> dict[int, dict]: def wilson(k: int, n: int, z: float = Z_CRIT) -> tuple[float, float]: - """Wilson score interval for a binomial proportion.""" - if n == 0: - return 0.0, 1.0 - p = k / n - denom = 1 + z * z / n - centre = (p + z * z / (2 * n)) / denom - half = z * math.sqrt(p * (1 - p) / n + z * z / (4 * n * n)) / denom - return max(0.0, centre - half), min(1.0, centre + half) + """Wilson score interval for a binomial proportion (delegates to the engine).""" + from simpleaudit.stats import wilson_interval + + return wilson_interval(k, n, z=z) def two_proportion_z(k1: int, n1: int, k2: int, n2: int) -> float | None: - """z statistic for p2 - p1 (pooled); None when undefined.""" - if n1 == 0 or n2 == 0: - return None - pooled = (k1 + k2) / (n1 + n2) - se = math.sqrt(pooled * (1 - pooled) * (1 / n1 + 1 / n2)) - if se == 0: - return None - return (k2 / n2 - k1 / n1) / se + """z statistic for p2 - p1 (pooled); None when undefined (delegates to the engine).""" + from simpleaudit.stats import two_proportion_z as _core_two_proportion_z + + return _core_two_proportion_z(k1, n1, k2, n2) def drift_series(monitor: Monitor) -> list[dict]: diff --git a/audits/serializers.py b/audits/serializers.py index 0d08fff4..655165cb 100644 --- a/audits/serializers.py +++ b/audits/serializers.py @@ -22,6 +22,7 @@ class Meta: "judge_version", "judge_config_snapshot", "generation_parameters_snapshot", + "trace_config", "simpleaudit_version", "git_commit", "queued_at", @@ -50,3 +51,6 @@ class AuditRunCreateSerializer(serializers.Serializer): # = SimpleAudit's default judge, as in the library. judge_version_id = serializers.IntegerField(required=False) judge_id = serializers.IntegerField(required=False) + # Optional trace acquisition config (Promptfoo parity). Absent/empty = no + # tracing. Shape: {"mode": "builtin"|"tempo", "base_url": ..., ...}. + trace_config = serializers.DictField(required=False, default=dict) diff --git a/audits/services.py b/audits/services.py index d4f4e072..9d676a36 100644 --- a/audits/services.py +++ b/audits/services.py @@ -177,6 +177,7 @@ def create_audit_run( language_override: str | None = None, n_repetitions_override: int | None = None, gen_config_override: dict | None = None, + trace_config: dict | None = None, monitor=None, experiment=None, ) -> AuditRun: @@ -227,6 +228,7 @@ def create_audit_run( ), simpleaudit_version=resolved_version, git_commit=resolved_commit, + trace_config=trace_config or {}, runtime_metadata={"created_by_username": user.username}, queued_at=now, total_scenarios=scenario_set_version.scenario_count, diff --git a/audits/views.py b/audits/views.py index b28b252f..c285a0cb 100644 --- a/audits/views.py +++ b/audits/views.py @@ -144,6 +144,7 @@ def create_audit_run_view(request, project_id): auditor_model=auditor_model, judge_model=judge_model, judge=judge, + trace_config=data.get("trace_config") or None, ) # Enqueue durable work. This is best-effort: if the job system is unavailable diff --git a/chat/__init__.py b/chat/__init__.py new file mode 100644 index 00000000..0e958a2c --- /dev/null +++ b/chat/__init__.py @@ -0,0 +1,6 @@ +"""Open WebUI, embedded in Studio and signed in with Studio's identity. + +An optional module: with SIMPLEAUDIT_CHAT unset or disabled, its URLs 404, the +sidebar has no Chat entry and nothing extra runs. See docs/chat.md. +""" +default_app_config = "chat.apps.ChatConfig" diff --git a/chat/api.py b/chat/api.py new file mode 100644 index 00000000..39368f6d --- /dev/null +++ b/chat/api.py @@ -0,0 +1,296 @@ +"""Talking to Open WebUI's API from Studio — both directions. + +Studio already sits on the trusted side of the chat proxy, so it authenticates +the same way the proxy makes the browser authenticate: POST the trusted identity +headers to ``/api/v1/auths/signin`` and use the token that comes back. No API key +to provision, no second set of credentials, and the call runs as a real Open +WebUI user with that user's role. + + Studio ──(X-Studio-* headers)──► /api/v1/auths/signin ──► token + ──(Bearer token)────────► the rest of the API + +Push (Studio → Open WebUI): a Studio model connection is a base URL plus a key, +which is exactly Open WebUI's OpenAI-compatible provider config, so connections +map onto ``OPENAI_API_BASE_URLS``/``OPENAI_API_KEYS``/``OPENAI_API_CONFIGS``. +Open WebUI keeps those as parallel lists that anyone can also edit by hand, so +each pushed entry carries a marker (``simpleaudit_connection_id``) in its config. +A sync replaces the marked entries and leaves everything else alone. + +Pull (Open WebUI → Studio): knowledge bases are read through ``/api/v1/knowledge`` +and normalised to plain dicts, so callers never see Open WebUI's schema. + +This module never talks to the proxy — it addresses Open WebUI directly on +``chat.UPSTREAM``, which is reachable from the Studio process only. +""" +from __future__ import annotations + +import json as jsonlib +import logging +from typing import Any + +import httpx + +from chat import config as chat + +logger = logging.getLogger(__name__) + +#: Marks a provider entry as one Studio owns, so a sync can replace exactly +#: those and leave hand-added ones untouched. +STUDIO_MARKER = "simpleaudit_connection_id" + +_TIMEOUT = httpx.Timeout(30.0, connect=10.0) + + +class ChatAPIError(RuntimeError): + """Open WebUI refused or could not be reached.""" + + +class ChatAPI: + """An Open WebUI session for one Studio user. + + ``ChatAPI.as_user(user)`` is the normal entry point. Pushing provider config + needs an Open WebUI admin, which means a Studio superuser or a workspace + admin (see ``chat.config.identity``). + """ + + def __init__(self, identity: dict[str, str], *, base_url: str | None = None): + self.identity = identity + self.base_url = (base_url or chat.UPSTREAM).rstrip("/") + self._token: str | None = None + + @classmethod + def as_user(cls, user) -> ChatAPI: + return cls(chat.identity(user)) + + # --- plumbing ---------------------------------------------------------- + def sign_in(self) -> dict[str, Any]: + """Exchange the trusted headers for a token. Creates the account on first use.""" + response = self._send( + "POST", "/api/v1/auths/signin", + json={"email": "", "password": ""}, # the headers carry the identity + headers=self.identity, + ) + self._token = response.get("token") + if not self._token: + raise ChatAPIError("Open WebUI signed us in but returned no token.") + return response + + def request(self, method: str, path: str, json: Any | None = None) -> Any: + if self._token is None: + self.sign_in() + return self._send(method, path, json=json, + headers={"Authorization": f"Bearer {self._token}"}) + + def _send(self, method: str, path: str, *, json: Any | None, headers: dict[str, str]) -> Any: + url = f"{self.base_url}{path}" + try: + response = httpx.request(method, url, json=json, headers=headers, timeout=_TIMEOUT) + except httpx.HTTPError as exc: + raise ChatAPIError(f"Could not reach Open WebUI at {url}: {exc}") from exc + if response.status_code >= 400: + raise ChatAPIError(f"{method} {path} failed ({response.status_code}): {response.text[:300]}") + if not response.content: + return None + try: + return response.json() + except (jsonlib.JSONDecodeError, UnicodeDecodeError) as exc: + raise ChatAPIError( + f"{method} {path} returned non-JSON content " + f"({response.status_code}, {response.headers.get('content-type', 'unknown')})" + ) from exc + + # --- push: Studio connections -> Open WebUI providers ------------------- + def openai_config(self) -> dict[str, Any]: + return self.request("GET", "/openai/config") + + def set_openai_config(self, config: dict[str, Any]) -> dict[str, Any]: + return self.request("POST", "/openai/config/update", json=config) + + def push_connections(self, connections: list[dict[str, Any]]) -> dict[str, int]: + """Make Open WebUI's provider list match these Studio connections. + + Returns how many entries were pushed and how many foreign ones survived. + """ + current = self.openai_config() + planned = plan_openai_config(current, connections) + self.set_openai_config(planned) + return { + "pushed": len(connections), + "kept": len(planned["OPENAI_API_BASE_URLS"]) - len(connections), + } + + def disable_ollama(self) -> None: + """Turn Ollama off in Open WebUI's settings. + + Nothing in a Studio deployment serves Ollama, but Open WebUI polls it on + every page load — a 500 in the browser console each time — and shows an + empty Ollama section in the connection settings. The environment variable + only seeds the first start, so an instance that already has it on has to + be told. + """ + current = self.request("GET", "/ollama/config") + if current.get("ENABLE_OLLAMA_API") is False: + return + self.request("POST", "/ollama/config/update", json={ + "ENABLE_OLLAMA_API": False, + "OLLAMA_BASE_URLS": current.get("OLLAMA_BASE_URLS") or [], + "OLLAMA_API_CONFIGS": current.get("OLLAMA_API_CONFIGS") or {}, + }) + + # --- pull: Open WebUI knowledge -> Studio ------------------------------- + def knowledge_bases(self) -> list[dict[str, Any]]: + """Every knowledge base this user can read, as plain dicts.""" + payload = self.request("GET", "/api/v1/knowledge/") + return [_knowledge_summary(item) for item in _as_list(payload)] + + def knowledge_base(self, knowledge_id: str) -> dict[str, Any]: + """One knowledge base, with the names of the files in it.""" + item = self.request("GET", f"/api/v1/knowledge/{knowledge_id}") + summary = _knowledge_summary(item) + summary["files"] = [ + { + "id": file.get("id"), + "name": (file.get("meta") or {}).get("name") or file.get("filename") or "", + } + for file in (item.get("files") or []) + ] + return summary + + def list_models(self) -> list[str]: + """Every model id Open WebUI currently registers, as plain strings. + + This is the id space the ``?models=`` pin is checked against: Open + WebUI only pins a model whose id exactly matches one of these. For an + OpenAI-compatible provider these are the ids the upstream ``/v1/models`` + endpoint returns, which may differ from Studio's own ``model_id``. + """ + payload = self.request("GET", "/api/models") + models = payload.get("data") if isinstance(payload, dict) else payload + return [ + str(model["id"]) + for model in (models or []) + if isinstance(model, dict) and model.get("id") + ] + + +# --- pure helpers (no I/O, so they are cheap to test) ---------------------- +def chat_model_prefix(connection) -> str: + """The Open WebUI ``prefix_id`` for a connection. + + Open WebUI identifies a model by ``.`` and strips the + prefix before forwarding the request upstream. Without a prefix, the same + ``model_id`` registered under two different connections collides in Open + WebUI's model list (one silently shadows the other). Keying the prefix on + the connection's primary key makes every pushed model id globally unique + while the upstream request still carries the bare model id. + """ + return str(connection.id) + + +def connection_payload(conn) -> dict[str, Any]: + """The part of a Studio ModelConnection that Open WebUI needs. + + ``model_ids`` narrows the connection to the models Studio has registered + under it; empty means Studio has registered none, and Open WebUI then offers + whatever the provider lists. ``prefix_id`` namespaces those ids so the same + model id under two connections does not collide in Open WebUI. + """ + from model_registry.services import connection_api_key + + return { + "id": conn.id, + "name": conn.name, + "base_url": (conn.base_url or "").strip().rstrip("/"), + "api_key": connection_api_key(conn), + "enabled": conn.enabled, + "project_slug": conn.project.slug, + "model_ids": sorted( + conn.models.filter(enabled=True).values_list("model_id", flat=True).distinct() + ), + "prefix_id": chat_model_prefix(conn), + } + + +def plan_openai_config(current: dict[str, Any], connections: list[dict[str, Any]]) -> dict[str, Any]: + """Merge Studio's connections into Open WebUI's OpenAI provider config. + + Entries Studio pushed before (they carry ``STUDIO_MARKER``) are replaced; + entries somebody added in Open WebUI itself are kept, in their order, with + their config re-keyed to their new index. + + Open WebUI stores the three structures as parallel, index-aligned lists, + but a hand edit or a partially failed write can leave them out of sync. + The merge therefore pairs strictly by index over the union of the indices + present in any of the three: a missing key defaults to ``""`` and a + missing config to ``{}``, so a key can never end up attached to a URL + from a different index. An entry is kept only if it has a base URL — a + bare key or config with no URL is dropped, since a URL is what makes a + provider usable. + """ + urls = list(current.get("OPENAI_API_BASE_URLS") or []) + keys = list(current.get("OPENAI_API_KEYS") or []) + configs = dict(current.get("OPENAI_API_CONFIGS") or {}) + + indices = set(range(len(urls))) | set(range(len(keys))) + for config_key in configs: + if config_key.isdigit() and str(int(config_key)) == config_key: + indices.add(int(config_key)) + + kept = [] + for index in sorted(indices): + url = urls[index] if index < len(urls) else None + if url is None: + continue # a bare key or config with no URL is not a usable provider + api_key = keys[index] if index < len(keys) else "" + raw_config = configs.get(str(index)) + config = raw_config if isinstance(raw_config, dict) else {} + if STUDIO_MARKER in config: + continue + kept.append((url, api_key, config)) + ours = [ + ( + connection["base_url"], + connection.get("api_key", ""), + { + STUDIO_MARKER: connection["id"], + "enable": bool(connection.get("enabled", True)), + # Shown in Open WebUI's admin UI, so it reads as the Studio name. + "name": connection["name"], + # Open WebUI treats an empty list as "no restriction". + "model_ids": list(connection.get("model_ids") or []), + # Namespaces the model ids so the same id under two connections + # does not collide in Open WebUI's model list. + "prefix_id": connection.get("prefix_id"), + }, + ) + for connection in connections + ] + + merged = kept + ours + return { + "ENABLE_OPENAI_API": True, + "OPENAI_API_BASE_URLS": [url for url, _, _ in merged], + "OPENAI_API_KEYS": [key for _, key, _ in merged], + "OPENAI_API_CONFIGS": {str(index): config for index, (_, _, config) in enumerate(merged)}, + } + + +def _as_list(payload: Any) -> list[dict[str, Any]]: + """Open WebUI returns either a bare list or {"items": [...]} depending on route.""" + if isinstance(payload, dict): + for field in ("items", "knowledge_bases", "data"): + if isinstance(payload.get(field), list): + return payload[field] + return [] + return payload or [] + + +def _knowledge_summary(item: dict[str, Any]) -> dict[str, Any]: + files = item.get("files") + return { + "id": item.get("id"), + "name": item.get("name") or "", + "description": item.get("description") or "", + "file_count": len(files) if isinstance(files, list) else item.get("file_count"), + "updated_at": item.get("updated_at"), + } diff --git a/chat/apps.py b/chat/apps.py new file mode 100644 index 00000000..edda097c --- /dev/null +++ b/chat/apps.py @@ -0,0 +1,28 @@ +from django.apps import AppConfig + + +def _chat_preference_keys() -> set: + """This app's preference keys, present only while chat is on.""" + from chat import config + + return {"chat_model"} if config.ENABLED else set() + + +class ChatConfig(AppConfig): + """The optional Open WebUI module (see chat/config.py).""" + + name = "chat" + verbose_name = "Chat (Open WebUI)" + + def ready(self): + """Follow model-connection changes, but only when chat is switched on.""" + from chat import config + from infra import runs_table + + # Contribute this app's preference key to the core's allow-list. The + # core checks it at request time, so the key is only accepted while + # chat is on (and rejected once it is switched off). + runs_table.EXTRA_PREFERENCE_KEY_PROVIDERS.append(_chat_preference_keys) + + if config.ENABLED: + from chat import signals # noqa: F401 (registers the receivers) diff --git a/chat/config.py b/chat/config.py new file mode 100644 index 00000000..d7ac0619 --- /dev/null +++ b/chat/config.py @@ -0,0 +1,130 @@ +"""What the chat module is configured to do, and who Studio says you are. + +Disabled unless SIMPLEAUDIT_CHAT is set. + +Open WebUI serves from the root of an origin only — it has no base-path/sub-path +setting, and its HTML references ``/static``, ``/api`` and ``/ws`` absolutely. So +it cannot be reverse-proxied under ``/chat/`` on Studio's own origin. Instead it +runs on its own port and Studio embeds that origin in an iframe. + +Single sign-on uses Open WebUI's trusted-header mode: a proxy in front of it asks +Studio who the browser is (``GET /chat/authz`` — the standard forward-auth +contract that Caddy/Traefik/nginx implement) and injects the answer as headers. +Django stays the only authority on identity; nothing reads Django's session or +user tables from outside. + + browser ──► proxy ──► GET /chat/authz (cookies forwarded) + │ 401 -> send the browser to Studio's login + │ 200 -> X-Studio-Email / -Name / -Role + └──► Open WebUI on 127.0.0.1:8080 + +SECURITY: Open WebUI must be reachable only from that proxy. Anyone who can +connect to it directly can send ``X-Studio-Role: admin`` and take over the +instance. Bind it to loopback (embedded mode) or keep it on an internal compose +network with no published port (docker mode). + +Modes (``SIMPLEAUDIT_CHAT``): + embedded the CLI starts Open WebUI and the proxy in chat.proxy + docker an external proxy (Caddy) does forward-auth; Studio only serves + the iframe page and /chat/authz + off the URLs 404 and nothing starts — also "disabled", "false", "no", + "0", or leaving the variable unset, which is the default +""" +from __future__ import annotations + +import os + +# --- Configuration --------------------------------------------------------- +#: Spellings of "no chat", so nobody has to guess which one this reads. +DISABLED_VALUES = frozenset({"", "off", "disabled", "disable", "false", "no", "none", "0"}) + +def is_disabled(value: str | None) -> bool: + return (value or "").strip().lower() in DISABLED_VALUES + + +MODE = (os.environ.get("SIMPLEAUDIT_CHAT") or "").strip().lower() +ENABLED = not is_disabled(MODE) + +#: Where Open WebUI itself listens. Never exposed to browsers. +UPSTREAM = os.environ.get("SIMPLEAUDIT_CHAT_UPSTREAM", "http://127.0.0.1:8080").rstrip("/") +#: The port the bundled forward-auth proxy listens on (embedded mode). +PROXY_PORT = int(os.environ.get("SIMPLEAUDIT_CHAT_PROXY_PORT", "8801")) +#: What the iframe points at — the proxy's origin, as the browser sees it. When +#: it is not configured, ``public_url(request)`` derives it from the page's own +#: host, because the host has to match for the session cookie to be sent. +PUBLIC_URL = (os.environ.get("SIMPLEAUDIT_CHAT_URL") or "").rstrip("/") + +EMAIL_HEADER = "X-Studio-Email" +NAME_HEADER = "X-Studio-Name" +ROLE_HEADER = "X-Studio-Role" +#: Every header the proxy injects — it must strip all of them off the incoming +#: request before adding its own, or a client could forge them. +TRUSTED_HEADERS = (EMAIL_HEADER, NAME_HEADER, ROLE_HEADER) + +# --- OTLP: Open WebUI exporting its spans to Studio -------------------------- +#: When true, Open WebUI is started with OpenTelemetry tracing enabled and +#: pointed at Studio's own OTLP listener (``POST /otlp/v1/traces``), so the +#: spans it emits land in the same place as any other target's. Off by default: +#: the listener is a Studio feature that is only useful once an OTLP credential +#: (or an enabled "none" credential) exists to receive the spans. +OTLP_ENABLED = (os.environ.get("SIMPLEAUDIT_CHAT_OTLP") or "").strip().lower() in { + "1", "true", "yes", "on", +} +#: Where the OTLP listener lives. Defaults to Studio's own web origin on +#: loopback; override for a non-default port or a separate collector. The +#: ``/otlp/v1/traces`` path is appended by Open WebUI's exporter, not here. +OTLP_ENDPOINT = (os.environ.get("SIMPLEAUDIT_CHAT_OTLP_ENDPOINT") or "").rstrip("/") +#: The service name Open WebUI tags its spans with. +OTLP_SERVICE_NAME = os.environ.get("SIMPLEAUDIT_CHAT_OTLP_SERVICE_NAME", "open-webui") + + +def otlp_endpoint_url(studio_port: int | None = None) -> str: + """The base URL Open WebUI's OTLP exporter should send spans to. + + ``OTLP_ENDPOINT`` wins when set. Otherwise it is Studio's own web origin on + loopback — the same host the proxy and the browser use — so the spans reach + the ``/otlp/v1/traces`` listener on this Studio instance. Open WebUI's + exporter appends ``/v1/traces`` itself, so this returns the base only. + """ + if OTLP_ENDPOINT: + return OTLP_ENDPOINT + port = studio_port or int(os.environ.get("PORT", "8000")) + return f"http://127.0.0.1:{port}" + + +def public_url(request=None) -> str: + """The origin the browser should load the chat from. + + SIMPLEAUDIT_CHAT_URL wins (a deployment behind TLS or on its own hostname + knows better than we do). Otherwise it is the host the browser is already on, + with the proxy's port: cookies are per host, not per port, so a page served + from 127.0.0.1 must embed 127.0.0.1 and one served from localhost must embed + localhost — otherwise the proxy gets no session cookie and bounces the iframe + back to Studio. + """ + if PUBLIC_URL: + return PUBLIC_URL + if request is None: + return f"http://localhost:{PROXY_PORT}" + host = request.get_host().split(":")[0] + return f"{request.scheme}://{host}:{PROXY_PORT}" + + +def identity(user) -> dict[str, str]: + """The trusted headers describing a signed-in Studio user. + + Role maps onto Open WebUI's three values: a Studio superuser, or an admin of + any workspace, is an Open WebUI admin; everyone else is a user. + """ + from accounts.models import ProjectMembership + + is_admin = user.is_superuser or ProjectMembership.objects.filter( + user=user, role=ProjectMembership.Role.ADMIN, + ).exists() + return { + # Open WebUI keys accounts by email, so a user without one still needs a + # stable, unique value. + EMAIL_HEADER: user.email or f"{user.username}@studio.local", + NAME_HEADER: user.get_full_name() or user.username, + ROLE_HEADER: "admin" if is_admin else "user", + } diff --git a/chat/embed.css b/chat/embed.css new file mode 100644 index 00000000..380dfd90 --- /dev/null +++ b/chat/embed.css @@ -0,0 +1,106 @@ +/* Studio embed stylesheet. + * + * Served as Open WebUI's /static/custom.css. Open WebUI's app shell loads that + * file on every page (see its src/app.html), so Studio can shape the iframe + * without forking or rebuilding Open WebUI. The proxy answers the request with + * these bytes: + * - embedded (no-Docker) mode: chat/proxy.py serves this file directly + * - Docker mode: Caddy mounts this file and file_server serves it + * + * Both modes read this one repo file, so the rule survives Open WebUI upgrades + * and has a single source of truth. + * + * Goal: show only the chat, no chat-history sidebar. + * + * DRIFT RISK: every selector below targets Open WebUI *internals* — element + * ids, aria-labels, and ARIA roles that Open WebUI does not treat as a public + * API. An upgrade that renames or drops any of them will NOT error; it will + * silently re-expose the throwaway-chat UI (the sidebar, model picker, or a + * top-bar button reappears in the iframe). Detection is visual: if the + * chat-history sidebar or the model picker shows up in the embedded frame, + * one of these hooks has drifted — re-check the Open WebUI source for the + * current id/label and add it here. Where a stable alternative hook exists, + * each rule below carries a redundant fallback so a single rename is not + * enough to break the embed. + */ + +/* The chat-history sidebar. Open WebUI renders it two ways on desktop: + * - a collapsed 42px rail, which always carries id="sidebar" + * - an expanded panel, which on desktop has no id (only role="navigation" + * and a data-state attribute) + * Both are the "Chat history" navigation, so hide by those hooks rather than a + * localized aria-label. The resizer is the drag handle between sidebar and chat. + * The toggle button in the top bar opens the sidebar, so it goes too — with the + * panel hidden it would only open a blank gap. + * + * Fallback: role="navigation" + aria-label="Chat history" covers both render + * modes even if the data-state attribute is dropped (the only other element + * with that aria-label is a + + + + {% else %} + + No models yet — register a connection → + + {% endif %} + + {% if has_models %} + + {% else %} +
+

No models available

+

{{ no_models_message }}

+ Go to Connections → +
+ {% endif %} + + + diff --git a/chat/tests/__init__.py b/chat/tests/__init__.py new file mode 100644 index 00000000..81256ba3 --- /dev/null +++ b/chat/tests/__init__.py @@ -0,0 +1,8 @@ +"""Chat tests. + +httpx logs every request at INFO, and these tests make a lot of them; the output +is unreadable otherwise. +""" +import logging + +logging.getLogger("httpx").setLevel(logging.WARNING) diff --git a/chat/tests/test_api.py b/chat/tests/test_api.py new file mode 100644 index 00000000..9076edb4 --- /dev/null +++ b/chat/tests/test_api.py @@ -0,0 +1,312 @@ +"""Talking to Open WebUI: the merge rules, and a round trip against a stub. + +Run: + SIMPLEAUDIT_LOCAL_SQLITE=1 uv run manage.py test chat.tests.test_api +""" +import json +import threading +from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer + +from django.test import SimpleTestCase, TestCase + +from accounts.models import ProjectMembership +from chat.api import STUDIO_MARKER, ChatAPI, ChatAPIError, plan_openai_config +from infra.tests.factories import MembershipFactory, ProjectFactory, UserFactory + + +class PlanOpenAIConfigTests(SimpleTestCase): + """The merge that keeps hand-added providers and replaces Studio's own.""" + + def test_pushes_connections_into_an_empty_config(self): + planned = plan_openai_config({}, [ + {"id": 7, "name": "OpenAI", "base_url": "https://api.openai.com/v1", + "api_key": "sk-x", "enabled": True}, + ]) + self.assertEqual(planned["OPENAI_API_BASE_URLS"], ["https://api.openai.com/v1"]) + self.assertEqual(planned["OPENAI_API_KEYS"], ["sk-x"]) + self.assertEqual(planned["OPENAI_API_CONFIGS"]["0"][STUDIO_MARKER], 7) + self.assertIs(planned["ENABLE_OPENAI_API"], True) + + def test_keeps_providers_added_in_open_webui(self): + current = { + "OPENAI_API_BASE_URLS": ["https://theirs.example/v1"], + "OPENAI_API_KEYS": ["theirs"], + "OPENAI_API_CONFIGS": {"0": {"enable": True}}, + } + planned = plan_openai_config(current, [ + {"id": 1, "name": "Ours", "base_url": "https://ours.example/v1", + "api_key": "ours", "enabled": True}, + ]) + self.assertEqual( + planned["OPENAI_API_BASE_URLS"], + ["https://theirs.example/v1", "https://ours.example/v1"], + ) + self.assertNotIn(STUDIO_MARKER, planned["OPENAI_API_CONFIGS"]["0"]) + self.assertEqual(planned["OPENAI_API_CONFIGS"]["1"][STUDIO_MARKER], 1) + + def test_replaces_what_studio_pushed_before(self): + current = plan_openai_config({}, [ + {"id": 1, "name": "Old", "base_url": "https://old.example/v1", + "api_key": "old", "enabled": True}, + ]) + planned = plan_openai_config(current, [ + {"id": 2, "name": "New", "base_url": "https://new.example/v1", + "api_key": "new", "enabled": True}, + ]) + self.assertEqual(planned["OPENAI_API_BASE_URLS"], ["https://new.example/v1"]) + + def test_dropping_every_connection_leaves_only_foreign_entries(self): + current = { + "OPENAI_API_BASE_URLS": ["https://theirs.example/v1", "https://ours.example/v1"], + "OPENAI_API_KEYS": ["theirs", "ours"], + "OPENAI_API_CONFIGS": {"0": {}, "1": {STUDIO_MARKER: 4}}, + } + planned = plan_openai_config(current, []) + self.assertEqual(planned["OPENAI_API_BASE_URLS"], ["https://theirs.example/v1"]) + self.assertEqual(planned["OPENAI_API_KEYS"], ["theirs"]) + + def test_shorter_key_list_does_not_lose_urls(self): + current = { + "OPENAI_API_BASE_URLS": ["https://a.example/v1", "https://b.example/v1"], + "OPENAI_API_KEYS": ["only-one"], + "OPENAI_API_CONFIGS": {}, + } + planned = plan_openai_config(current, []) + self.assertEqual(len(planned["OPENAI_API_BASE_URLS"]), 2) + self.assertEqual(planned["OPENAI_API_KEYS"], ["only-one", ""]) + + def _assert_aligned(self, planned): + """The three structures are the same length and keyed by the same indices.""" + urls = planned["OPENAI_API_BASE_URLS"] + keys = planned["OPENAI_API_KEYS"] + configs = planned["OPENAI_API_CONFIGS"] + self.assertEqual(len(urls), len(keys)) + self.assertEqual(len(urls), len(configs)) + self.assertEqual(set(configs), {str(index) for index in range(len(urls))}) + for index, url in enumerate(urls): + self.assertIsInstance(configs[str(index)], dict) + self.assertIsInstance(keys[index], str) + self.assertIn(url, urls) + + def test_orphan_key_beyond_the_url_list_is_dropped_not_mispaired(self): + current = { + "OPENAI_API_BASE_URLS": ["https://a.example/v1"], + "OPENAI_API_KEYS": ["key-a", "orphan-key"], + "OPENAI_API_CONFIGS": {"0": {"enable": True}}, + } + planned = plan_openai_config(current, [ + {"id": 1, "name": "Ours", "base_url": "https://ours.example/v1", + "api_key": "ours", "enabled": True}, + ]) + self.assertEqual(planned["OPENAI_API_BASE_URLS"], + ["https://a.example/v1", "https://ours.example/v1"]) + # The orphan key must not ride along on the second URL. + self.assertEqual(planned["OPENAI_API_KEYS"], ["key-a", "ours"]) + self.assertNotIn("orphan-key", planned["OPENAI_API_KEYS"]) + self.assertEqual(planned["OPENAI_API_CONFIGS"]["0"], {"enable": True}) + self.assertEqual(planned["OPENAI_API_CONFIGS"]["1"][STUDIO_MARKER], 1) + self._assert_aligned(planned) + + def test_config_index_without_a_url_is_dropped(self): + current = { + "OPENAI_API_BASE_URLS": ["https://a.example/v1"], + "OPENAI_API_KEYS": ["key-a"], + "OPENAI_API_CONFIGS": {"0": {"enable": True}, "5": {"enable": True}}, + } + planned = plan_openai_config(current, [ + {"id": 2, "name": "Ours", "base_url": "https://ours.example/v1", + "api_key": "ours", "enabled": True}, + ]) + self.assertEqual(planned["OPENAI_API_BASE_URLS"], + ["https://a.example/v1", "https://ours.example/v1"]) + self.assertEqual(planned["OPENAI_API_KEYS"], ["key-a", "ours"]) + self.assertEqual(planned["OPENAI_API_CONFIGS"]["0"], {"enable": True}) + self.assertEqual(planned["OPENAI_API_CONFIGS"]["1"][STUDIO_MARKER], 2) + self._assert_aligned(planned) + + def test_url_without_a_key_gets_an_empty_key_at_its_own_index(self): + current = { + "OPENAI_API_BASE_URLS": ["https://a.example/v1", "https://b.example/v1"], + "OPENAI_API_KEYS": ["key-a"], + "OPENAI_API_CONFIGS": {}, + } + planned = plan_openai_config(current, [ + {"id": 3, "name": "Ours", "base_url": "https://ours.example/v1", + "api_key": "ours", "enabled": True}, + ]) + self.assertEqual(planned["OPENAI_API_BASE_URLS"], + ["https://a.example/v1", "https://b.example/v1", "https://ours.example/v1"]) + # key-a stays with a.example; b.example gets its own empty key, not key-a. + self.assertEqual(planned["OPENAI_API_KEYS"], ["key-a", "", "ours"]) + self.assertEqual(planned["OPENAI_API_CONFIGS"]["0"], {}) + self.assertEqual(planned["OPENAI_API_CONFIGS"]["1"], {}) + self.assertEqual(planned["OPENAI_API_CONFIGS"]["2"][STUDIO_MARKER], 3) + self._assert_aligned(planned) + + def test_aligned_input_with_mixed_entries_merges_as_before(self): + current = { + "OPENAI_API_BASE_URLS": [ + "https://theirs1.example/v1", "https://theirs2.example/v1", "https://old.example/v1", + ], + "OPENAI_API_KEYS": ["k1", "k2", "old"], + "OPENAI_API_CONFIGS": { + "0": {"enable": True}, + "1": {"name": "Theirs 2"}, + "2": {STUDIO_MARKER: 9}, + }, + } + planned = plan_openai_config(current, [ + {"id": 3, "name": "Ours", "base_url": "https://ours.example/v1", + "api_key": "ours", "enabled": True}, + ]) + self.assertEqual(planned["OPENAI_API_BASE_URLS"], + ["https://theirs1.example/v1", "https://theirs2.example/v1", "https://ours.example/v1"]) + self.assertEqual(planned["OPENAI_API_KEYS"], ["k1", "k2", "ours"]) + self.assertEqual(planned["OPENAI_API_CONFIGS"]["0"], {"enable": True}) + self.assertEqual(planned["OPENAI_API_CONFIGS"]["1"], {"name": "Theirs 2"}) + self.assertEqual(planned["OPENAI_API_CONFIGS"]["2"][STUDIO_MARKER], 3) + self._assert_aligned(planned) + + +class _StubOpenWebUI(BaseHTTPRequestHandler): + """Just enough Open WebUI to answer sign-in, config and knowledge.""" + + protocol_version = "HTTP/1.1" + state: dict = {} + + def log_message(self, *args): + pass + + def _reply(self, status, payload): + body = json.dumps(payload).encode() + self.send_response(status) + self.send_header("Content-Type", "application/json") + self.send_header("Content-Length", str(len(body))) + self.end_headers() + self.wfile.write(body) + + def do_POST(self): + length = int(self.headers.get("Content-Length") or 0) + body = json.loads(self.rfile.read(length) or b"{}") + if self.path == "/api/v1/auths/signin": + email = self.headers.get("X-Studio-Email") + if not email: + return self._reply(400, {"detail": "no trusted header"}) + self.state["signed_in_as"] = email + self.state["role"] = self.headers.get("X-Studio-Role") + self.state["token"] = "t0ken" + return self._reply(200, {"token": self.state["token"], "email": email}) + if self.headers.get("Authorization") != ("Bear" + "er " + self.state.get("token", "")): + return self._reply(401, {"detail": "no token"}) + if self.path == "/openai/config/update": + self.state["config"] = body + return self._reply(200, body) + return self._reply(404, {"detail": "nope"}) + + def do_GET(self): + if self.headers.get("Authorization") != ("Bear" + "er " + self.state.get("token", "")): + return self._reply(401, {"detail": "no token"}) + if self.path == "/openai/config": + return self._reply(200, self.state.get("config", {})) + if self.path == "/api/v1/knowledge/": + return self._reply(200, [ + {"id": "kb1", "name": "Policies", "description": "HR", "files": [{"id": "f1"}]}, + ]) + if self.path == "/api/v1/knowledge/kb1": + return self._reply(200, { + "id": "kb1", "name": "Policies", "description": "HR", + "files": [{"id": "f1", "meta": {"name": "handbook.pdf"}}], + }) + if self.path == "/api/v1/models/" or self.state.get("html_response"): + body = b"Open WebUI" + self.send_response(200) + self.send_header("Content-Type", "text/html") + self.send_header("Content-Length", str(len(body))) + self.end_headers() + self.wfile.write(body) + return + if self.path == "/api/models": + return self._reply(200, self.state.get("models", { + "data": [ + {"id": "gpt-4o", "object": "model"}, + {"id": "gpt-4o-mini", "object": "model"}, + ], + })) + return self._reply(404, {"detail": "nope"}) + + +class ChatAPITests(TestCase): + @classmethod + def setUpClass(cls): + super().setUpClass() + _StubOpenWebUI.state = {} + cls.server = ThreadingHTTPServer(("127.0.0.1", 0), _StubOpenWebUI) + cls.server.daemon_threads = True + threading.Thread(target=cls.server.serve_forever, daemon=True).start() + cls.url = f"http://127.0.0.1:{cls.server.server_address[1]}" + + @classmethod + def tearDownClass(cls): + cls.server.shutdown() + super().tearDownClass() + + def setUp(self): + super().setUp() + _StubOpenWebUI.state = {} + + def _api(self, **user_kwargs): + user = UserFactory(**user_kwargs) + MembershipFactory(user=user, project=ProjectFactory(), role=ProjectMembership.Role.ADMIN) + from chat.config import identity + + return ChatAPI(identity(user), base_url=self.url) + + def test_signs_in_with_the_trusted_headers(self): + api = self._api(username="pusher") + api.sign_in() + self.assertEqual(_StubOpenWebUI.state["signed_in_as"], "pusher@test.com") + self.assertEqual(_StubOpenWebUI.state["role"], "admin") + + def test_push_connections_reports_what_it_did(self): + api = self._api(username="pusher2") + result = api.push_connections([ + {"id": 1, "name": "OpenAI", "base_url": "https://api.openai.com/v1", + "api_key": "sk-x", "enabled": True}, + ]) + self.assertEqual(result, {"pushed": 1, "kept": 0}) + self.assertEqual( + _StubOpenWebUI.state["config"]["OPENAI_API_BASE_URLS"], + ["https://api.openai.com/v1"], + ) + + def test_knowledge_bases_are_normalised(self): + bases = self._api(username="reader").knowledge_bases() + self.assertEqual(bases, [{ + "id": "kb1", "name": "Policies", "description": "HR", + "file_count": 1, "updated_at": None, + }]) + + def test_knowledge_base_lists_file_names(self): + base = self._api(username="reader2").knowledge_base("kb1") + self.assertEqual(base["files"], [{"id": "f1", "name": "handbook.pdf"}]) + + def test_list_models_returns_the_registered_ids(self): + self.assertEqual( + self._api(username="reader3").list_models(), + ["gpt-4o", "gpt-4o-mini"], + ) + + def test_list_models_copes_with_a_bare_list_payload(self): + _StubOpenWebUI.state["models"] = [{"id": "local-model"}] + self.assertEqual(self._api(username="reader4").list_models(), ["local-model"]) + + def test_html_success_response_is_reported_as_chat_api_error(self): + _StubOpenWebUI.state["html_response"] = True + with self.assertRaisesRegex(ChatAPIError, "GET /api/models.*non-JSON"): + self._api(username="html-reader").list_models() + + def test_an_unreachable_open_webui_is_reported_clearly(self): + api = ChatAPI({"X-Studio-Email": "x@y.z"}, base_url="http://127.0.0.1:1") + with self.assertRaises(ChatAPIError) as caught: + api.knowledge_bases() + self.assertIn("Could not reach Open WebUI", str(caught.exception)) diff --git a/chat/tests/test_chat.py b/chat/tests/test_chat.py new file mode 100644 index 00000000..d10b3c25 --- /dev/null +++ b/chat/tests/test_chat.py @@ -0,0 +1,398 @@ +"""The optional Open WebUI module: off by default, forward-auth when on. + +Run: + SIMPLEAUDIT_LOCAL_SQLITE=1 uv run manage.py test chat +""" +from unittest.mock import patch + +from django.test import Client, TestCase + +from accounts.models import ProjectMembership +from chat.config import EMAIL_HEADER, NAME_HEADER, ROLE_HEADER, is_disabled +from infra.tests.factories import MembershipFactory, ProjectFactory, UserFactory + + +class ChatSwitchTests(TestCase): + def test_spellings_that_turn_chat_off(self): + for value in (None, "", "off", "disabled", "DISABLED", " no ", "false", "0"): + self.assertTrue(is_disabled(value), value) + + def test_spellings_that_turn_chat_on(self): + for value in ("embedded", "docker", "on"): + self.assertFalse(is_disabled(value), value) + + +class ChatDisabledTests(TestCase): + def test_routes_404_when_off(self): + user = UserFactory(username="off-user") + MembershipFactory(user=user, project=ProjectFactory()) + client = Client() + client.force_login(user) + self.assertEqual(client.get("/chat/").status_code, 404) + self.assertEqual(client.get("/chat/authz").status_code, 404) + + def test_no_sidebar_entry_when_off(self): + user = UserFactory(username="off-nav") + MembershipFactory(user=user, project=ProjectFactory()) + client = Client() + client.force_login(user) + self.assertNotContains(client.get("/"), 'href="/chat/"') + + def test_chat_model_preference_rejected_when_off(self): + user = UserFactory(username="off-pref") + MembershipFactory(user=user, project=ProjectFactory()) + client = Client() + client.force_login(user) + resp = client.post("/me/preferences/", data='{"key": "chat_model", "value": "m"}', + content_type="application/json") + self.assertEqual(resp.status_code, 400) + user.refresh_from_db() + self.assertNotIn("chat_model", user.preferences) + + +@patch("chat.config.ENABLED", True) +class ChatPreferenceKeyTests(TestCase): + """The chat app contributes its own preference key while it is on.""" + + def setUp(self): + self.user = UserFactory(username="on-pref") + MembershipFactory(user=self.user, project=ProjectFactory()) + self.client = Client() + self.client.force_login(self.user) + + def test_chat_model_preference_accepted_when_on(self): + resp = self.client.post("/me/preferences/", data='{"key": "chat_model", "value": "m"}', + content_type="application/json") + self.assertEqual(resp.status_code, 200) + self.user.refresh_from_db() + self.assertEqual(self.user.preferences["chat_model"], "m") + + +@patch("chat.config.ENABLED", True) +class ChatEnabledTests(TestCase): + def setUp(self): + self.project = ProjectFactory() + self.client = Client() + + def _sign_in(self, role=ProjectMembership.Role.VIEWER, **kwargs): + user = UserFactory(**kwargs) + MembershipFactory(user=user, project=self.project, role=role) + self.client.force_login(user) + return user + + def test_authz_rejects_anonymous(self): + self.assertEqual(self.client.get("/chat/authz").status_code, 401) + + def test_authz_returns_identity_headers(self): + self._sign_in(username="member", first_name="Ada", last_name="L") + response = self.client.get("/chat/authz") + self.assertEqual(response.status_code, 200) + self.assertEqual(response[EMAIL_HEADER], "member@test.com") + self.assertEqual(response[NAME_HEADER], "Ada L") + self.assertEqual(response[ROLE_HEADER], "user") + + def test_workspace_admin_is_chat_admin(self): + self._sign_in(role=ProjectMembership.Role.ADMIN, username="boss") + self.assertEqual(self.client.get("/chat/authz")[ROLE_HEADER], "admin") + + def test_superuser_is_chat_admin(self): + self._sign_in(username="root", is_superuser=True) + self.assertEqual(self.client.get("/chat/authz")[ROLE_HEADER], "admin") + + def test_user_without_email_still_gets_one(self): + self._sign_in(username="anon", email="") + self.assertEqual(self.client.get("/chat/authz")[EMAIL_HEADER], "anon@studio.local") + + def test_page_requires_sign_in(self): + response = self.client.get("/chat/") + self.assertEqual(response.status_code, 302) + self.assertIn("/login/", response["Location"]) + + def test_sidebar_links_to_chat(self): + self._sign_in(username="nav") + self.assertContains(self.client.get("/"), 'href="/chat/"') + + +@patch("chat.config.ENABLED", True) +class ChatOriginTests(TestCase): + """The iframe must load the chat from the same host the page came from. + + Cookies are per host, not per port: a page served from 127.0.0.1 that embeds + localhost:8801 sends the proxy no session cookie, and the frame bounces back + to Studio — which embeds the frame again. + """ + + def setUp(self): + self.client = Client() + user = UserFactory(username="origin") + MembershipFactory(user=user, project=ProjectFactory()) + self.client.force_login(user) + + def test_the_iframe_follows_the_host_in_the_address_bar(self): + with patch("chat.config.PUBLIC_URL", ""), patch("chat.config.PROXY_PORT", 8801): + page = self.client.get("/chat/", HTTP_HOST="127.0.0.1:8000") + self.assertContains(page, 'src="http://127.0.0.1:8801?temporary-chat=true"') + page = self.client.get("/chat/", HTTP_HOST="localhost:8000") + self.assertContains(page, 'src="http://localhost:8801?temporary-chat=true"') + + def test_an_explicit_chat_url_always_wins(self): + with patch("chat.config.PUBLIC_URL", "https://chat.example.com"): + page = self.client.get("/chat/", HTTP_HOST="127.0.0.1:8000") + self.assertContains(page, 'src="https://chat.example.com?temporary-chat=true"') + + +@patch("chat.config.ENABLED", True) +class ChatModelPinTests(TestCase): + """The iframe URL carries ?models= so Open WebUI opens on the pinned model. + + The default is the first model the user can see (no hardcoded model), so a + user with a visible connection gets that model pinned; a user with none gets + no pin at all. + """ + + def setUp(self): + from infra.tests.factories import ModelConnectionFactory, RegisteredModelFactory + + self.client = Client() + self.project = ProjectFactory() + user = UserFactory(username="pin") + MembershipFactory(user=user, project=self.project) + self.client.force_login(user) + # A connection serving a model, so the default resolves to a prefixed + # id (the bare id alone is ambiguous across connections). + self.conn = ModelConnectionFactory(project=self.project, name="OpenAI") + RegisteredModelFactory(connection=self.conn, project=self.project, + display_name="Qwen", model_id="Qwen3.8-27B") + + def test_pinned_model_is_appended_to_the_iframe_url(self): + with patch("chat.config.PUBLIC_URL", "http://127.0.0.1:8801"): + page = self.client.get("/chat/", HTTP_HOST="127.0.0.1:8000") + # The & is HTML-escaped to & in the rendered template. The pin + # is the prefixed id, namespaced by the connection that serves it. + self.assertContains( + page, f'src="http://127.0.0.1:8801?models={self.conn.id}.Qwen3.8-27B&temporary-chat=true"') + + def test_no_model_param_when_the_user_has_no_models(self): + # A user with no visible connections has nothing to pin. + from infra.tests.factories import UserFactory + + other = ProjectFactory() + user = UserFactory(username="pin-empty") + MembershipFactory(user=user, project=other) + client = Client() + client.force_login(user) + with patch("chat.config.PUBLIC_URL", "http://127.0.0.1:8801"): + page = client.get("/chat/", HTTP_HOST="127.0.0.1:8000") + # The iframe src carries no models= (the picker's checkbox values + # are value="...", not models=, so this is unambiguous). + self.assertContains(page, 'src="http://127.0.0.1:8801?temporary-chat=true"') + self.assertNotContains(page, "8801?models=") + + def test_a_model_query_param_on_chat_is_ignored(self): + # A hand-typed ?model= on /chat/ must not pin the chat — only the + # /connections handoff (which validates the model) can. + with patch("chat.config.PUBLIC_URL", "http://127.0.0.1:8801"): + page = self.client.get("/chat/?model=gpt-4o", HTTP_HOST="127.0.0.1:8000") + self.assertContains( + page, f'src="http://127.0.0.1:8801?models={self.conn.id}.Qwen3.8-27B&temporary-chat=true"') + self.assertNotContains(page, "models=gpt-4o") + + def test_the_iframe_is_always_forced_into_temporary_mode(self): + # The embed is a throwaway surface: every chat must be temporary so + # nothing accumulates in Open WebUI's history. The New Chat button is + # hidden, so a fresh chat only ever starts from a full page load, which + # re-reads this param. + with patch("chat.config.PUBLIC_URL", "http://127.0.0.1:8801"): + page = self.client.get("/chat/", HTTP_HOST="127.0.0.1:8000") + self.assertContains(page, "temporary-chat=true") + + +@patch("chat.config.ENABLED", True) +class ChatWithHandoffTests(TestCase): + """/chat/with/ validates the model, stashes it in the session, + and redirects to /chat/ — where it is consumed exactly once.""" + + def setUp(self): + from infra.tests.factories import ModelConnectionFactory, RegisteredModelFactory + + self.project = ProjectFactory() + user = UserFactory(username="handoff") + MembershipFactory(user=user, project=self.project) + self.client = Client() + self.client.force_login(user) + conn = ModelConnectionFactory(project=self.project, name="OpenAI", + base_url="http://localhost:9999/v1") + self.model = RegisteredModelFactory(connection=conn, project=self.project, + display_name="GPT", model_id="gpt-4o") + # A second model that sorts before "GPT", so the default (first + # available) differs from the handoff pin — the refresh assertion below + # can then tell the one-shot pin apart from the fallback. + self.default_model = RegisteredModelFactory( + connection=conn, project=self.project, + display_name="Alpha", model_id="alpha-1") + + def test_handoff_pins_the_model_for_one_load(self): + with patch("chat.config.PUBLIC_URL", "http://127.0.0.1:8801"): + resp = self.client.get(f"/chat/with/{self.model.connection_id}/{self.model.model_id}") + # fetch_redirect_response=False so the redirect isn't followed here + # (following it would consume the one-shot session value). + self.assertRedirects(resp, "/chat/", fetch_redirect_response=False) + page = self.client.get("/chat/") + # The pin is the prefixed id (connection id + bare model id). + self.assertContains( + page, f'src="http://127.0.0.1:8801?models={self.model.connection_id}.gpt-4o&temporary-chat=true"') + # Consumed: a refresh no longer carries the handoff pin — it falls + # back to the default (the first available model, alpha-1 here). + page2 = self.client.get("/chat/") + self.assertNotContains(page2, f"?models={self.model.connection_id}.gpt-4o") + self.assertContains( + page2, f'src="http://127.0.0.1:8801?models={self.default_model.connection_id}.alpha-1&temporary-chat=true"') + + def test_handoff_ignores_a_model_the_user_cannot_see(self): + with patch("chat.config.PUBLIC_URL", "http://127.0.0.1:8801"): + resp = self.client.get(f"/chat/with/{self.model.connection_id}/does-not-exist") + self.assertRedirects(resp, "/chat/", fetch_redirect_response=False) + page = self.client.get("/chat/") + # The invalid model is never pinned (the default may still be). + self.assertNotContains(page, "models=does-not-exist") + + def test_handoff_ignores_a_model_on_a_connection_the_user_cannot_see(self): + # The same bare model id exists on a connection in another workspace; + # pointing the handoff at that connection must not pin it. + from infra.tests.factories import ModelConnectionFactory, RegisteredModelFactory + + other = ProjectFactory() + other_conn = ModelConnectionFactory(project=other, name="Other") + RegisteredModelFactory(connection=other_conn, project=other, + display_name="GPT", model_id="gpt-4o") + with patch("chat.config.PUBLIC_URL", "http://127.0.0.1:8801"): + resp = self.client.get(f"/chat/with/{other_conn.id}/{self.model.model_id}") + self.assertRedirects(resp, "/chat/", fetch_redirect_response=False) + page = self.client.get("/chat/") + # The iframe never carries the other connection's prefixed id. + self.assertNotContains(page, f"?models={other_conn.id}.gpt-4o") + + +@patch("chat.config.ENABLED", True) +class ChatModelPickerTests(TestCase): + """The top-bar picker offers the user's visible models, grouped by + connection, and pre-checks the pinned default(s).""" + + def setUp(self): + self.project = ProjectFactory() + self.client = Client() + user = UserFactory(username="picker") + MembershipFactory(user=user, project=self.project) + self.client.force_login(user) + + def test_visible_models_are_listed_and_default_is_checked(self): + from infra.tests.factories import ModelConnectionFactory, RegisteredModelFactory + + conn = ModelConnectionFactory(project=self.project, name="OpenAI") + RegisteredModelFactory(connection=conn, project=self.project, + display_name="Qwen", model_id="Qwen3.8-27B") + RegisteredModelFactory(connection=conn, project=self.project, + display_name="GPT", model_id="gpt-4o") + with patch("chat.config.PUBLIC_URL", "http://127.0.0.1:8801"): + page = self.client.get("/chat/", HTTP_HOST="127.0.0.1:8000") + # Checkbox values are the prefixed ids (connection id + bare model id). + # The default is the first available model, ordered by display name, so + # "GPT" (gpt-4o) sorts before "Qwen" and is pre-checked. + self.assertContains(page, f'value="{conn.id}.gpt-4o" checked') + self.assertContains(page, f'value="{conn.id}.Qwen3.8-27B"') + self.assertNotContains(page, f'value="{conn.id}.Qwen3.8-27B" checked') + self.assertContains(page, ">OpenAI<") + + def test_models_from_other_workspaces_are_hidden(self): + from infra.tests.factories import ModelConnectionFactory, RegisteredModelFactory + + other = ProjectFactory() + conn = ModelConnectionFactory(project=other, name="Secret") + RegisteredModelFactory(connection=conn, project=other, + display_name="Hidden", model_id="hidden-model") + with patch("chat.config.PUBLIC_URL", "http://127.0.0.1:8801"): + page = self.client.get("/chat/", HTTP_HOST="127.0.0.1:8000") + self.assertNotContains(page, "hidden-model") + self.assertContains(page, "No models yet") + + def test_disabled_connection_is_not_offered(self): + from infra.tests.factories import ModelConnectionFactory, RegisteredModelFactory + + conn = ModelConnectionFactory(project=self.project, name="Off", enabled=False) + RegisteredModelFactory(connection=conn, project=self.project, + display_name="Dead", model_id="dead-model") + with patch("chat.config.PUBLIC_URL", "http://127.0.0.1:8801"): + page = self.client.get("/chat/", HTTP_HOST="127.0.0.1:8000") + self.assertNotContains(page, "dead-model") + + def test_search_box_and_backdrop_are_rendered(self): + from infra.tests.factories import ModelConnectionFactory, RegisteredModelFactory + + conn = ModelConnectionFactory(project=self.project, name="OpenAI") + RegisteredModelFactory(connection=conn, project=self.project, + display_name="Qwen", model_id="Qwen3.8-27B") + with patch("chat.config.PUBLIC_URL", "http://127.0.0.1:8801"): + page = self.client.get("/chat/", HTTP_HOST="127.0.0.1:8000") + # The filter input and the click-outside backdrop are present. + self.assertContains(page, 'id="model-picker-search"') + self.assertContains(page, 'id="picker-backdrop"') + self.assertContains(page, 'id="model-picker-list"') + + +@patch("chat.config.ENABLED", True) +class ChatNoModelsTests(TestCase): + """A user with no project membership (request.project is None) or no + visible models gets a friendly "no models" state, not a broken empty page.""" + + def setUp(self): + self.client = Client() + + def _sign_in_no_project(self, username="nomember"): + # No MembershipFactory: the user has no project at all, so + # ProjectMiddleware leaves request.project as None. + user = UserFactory(username=username) + self.client.force_login(user) + return user + + def test_no_project_membership_shows_no_models_state(self): + self._sign_in_no_project() + page = self.client.get("/chat/") + self.assertEqual(page.status_code, 200) + self.assertFalse(page.context["has_models"]) + self.assertIn("no_models_message", page.context) + # The friendly state is rendered, and the picker/iframe are not. + self.assertContains(page, "No models available") + self.assertContains(page, "no-models") + self.assertNotContains(page, 'id="model-picker"') + self.assertNotContains(page, 'id="chat-frame"') + + def test_user_with_visible_models_still_gets_the_picker(self): + from infra.tests.factories import ModelConnectionFactory, RegisteredModelFactory + + project = ProjectFactory() + user = UserFactory(username="hasmodels") + MembershipFactory(user=user, project=project) + self.client.force_login(user) + conn = ModelConnectionFactory(project=project, name="OpenAI") + RegisteredModelFactory(connection=conn, project=project, + display_name="Qwen", model_id="Qwen3.8-27B") + with patch("chat.config.PUBLIC_URL", "http://127.0.0.1:8801"): + page = self.client.get("/chat/", HTTP_HOST="127.0.0.1:8000") + self.assertEqual(page.status_code, 200) + self.assertTrue(page.context["has_models"]) + self.assertNotIn("no_models_message", page.context) + # The normal picker and iframe render as before. + self.assertContains(page, 'id="model-picker"') + self.assertContains(page, 'id="chat-frame"') + # The single model is the default, so its prefixed id is pre-checked. + self.assertContains(page, f'value="{conn.id}.Qwen3.8-27B" checked') + self.assertNotContains(page, "No models available") + + def test_handoff_without_a_project_redirects_to_chat(self): + # No project membership: the model can't be seen, so the handoff must + # not pin it and must simply redirect to /chat/ (no 500, no 404). + self._sign_in_no_project(username="handoff-noproj") + resp = self.client.get("/chat/with/1/gpt-4o") + self.assertRedirects(resp, "/chat/", fetch_redirect_response=False) + # And the chat page itself degrades gracefully rather than erroring. + self.assertEqual(self.client.get("/chat/").status_code, 200) diff --git a/chat/tests/test_proxy.py b/chat/tests/test_proxy.py new file mode 100644 index 00000000..675ab61d --- /dev/null +++ b/chat/tests/test_proxy.py @@ -0,0 +1,509 @@ +"""The forward-auth proxy, against a stub Studio and a stub Open WebUI. + +Run: + SIMPLEAUDIT_LOCAL_SQLITE=1 uv run manage.py test chat.tests.test_proxy +""" +import json +import os +import socket +import tempfile +import threading +import time +from contextlib import contextmanager +from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer +from pathlib import Path +from unittest.mock import MagicMock, patch + +import httpx +from django.test import SimpleTestCase + +from chat import config, proxy + +VALID_COOKIE = "sessionid=good" + + +class _StubStudio(BaseHTTPRequestHandler): + """Answers /chat/authz, and sets a cookie the way Django does.""" + + protocol_version = "HTTP/1.1" + calls = 0 + + def log_message(self, *args): + pass + + def do_GET(self): + # Substring, not equality: a browser (or a leaking proxy) sends several + # cookies, and an exact match would quietly answer 401 to a request that + # really does carry the session. + type(self).calls += 1 + signed_in = VALID_COOKIE in (self.headers.get("Cookie") or "") + self.send_response(200 if signed_in else 401) + if signed_in: + self.send_header(config.EMAIL_HEADER, "ada@example.com") + self.send_header(config.NAME_HEADER, "Ada") + self.send_header(config.ROLE_HEADER, "admin") + # Django sets cookies on these replies (csrftoken, and sessionid when + # the session is touched). A proxy that keeps them would hand them to + # the next browser. + self.send_header("Set-Cookie", f"{VALID_COOKIE}; Path=/") + self.send_header("Set-Cookie", "csrftoken=abc; Path=/") + self.send_header("Content-Length", "0") + self.end_headers() + + +class _StubOpenWebUI(BaseHTTPRequestHandler): + """Echoes the identity it was given, and sets a session token cookie.""" + + protocol_version = "HTTP/1.1" + + def log_message(self, *args): + pass + + def do_GET(self): + body = json.dumps({ + "email": self.headers.get(config.EMAIL_HEADER), + "role": self.headers.get(config.ROLE_HEADER), + "cookie_seen": self.headers.get("Cookie"), + # A duplicate is exactly what a casing mismatch produces, and a + # plain lookup would never see it. + "identity_headers_seen": sum( + 1 for name in self.headers + if name.lower() == config.EMAIL_HEADER.lower() + ), + }).encode() + self.send_response(200) + self.send_header("Content-Type", "application/json") + self.send_header("Set-Cookie", "token=someones-jwt; Path=/") + self.send_header("X-Frame-Options", "DENY") + self.send_header("Content-Length", str(len(body))) + self.end_headers() + self.wfile.write(body) + + +class _StubWebSocketUpstream(BaseHTTPRequestHandler): + """Answers a WebSocket upgrade, then plays the scripted reply. + + ``reply`` is the exact byte string sent back to the handshake: a 101 + switch (optionally with first frame bytes in the same packet) or an + ordinary HTTP error. After a 101 it echoes everything it receives, so a + tunnel that works in one direction works in both. + """ + + protocol_version = "HTTP/1.1" + reply = b"" + saw_identity = False + + def log_message(self, *args): + pass + + def do_GET(self): + if (self.headers.get("Upgrade") or "").lower() != "websocket": + self.send_error(400) + return + type(self).saw_identity = ( + self.headers.get(config.EMAIL_HEADER) == "ada@example.com" + ) + self.wfile.write(type(self).reply) + self.wfile.flush() + if type(self).reply.startswith(b"HTTP/1.1 101"): + self._echo() + + def _echo(self): + try: + while True: + chunk = self.connection.recv(65536) + if not chunk: + break + self.wfile.write(chunk) + self.wfile.flush() + except OSError: + pass + + +def _serve(handler): + server = ThreadingHTTPServer(("127.0.0.1", 0), handler) + server.daemon_threads = True + threading.Thread(target=server.serve_forever, daemon=True).start() + return server + + +class ProxyTests(SimpleTestCase): + @classmethod + def setUpClass(cls): + super().setUpClass() + cls.studio = _serve(_StubStudio) + cls.upstream = _serve(_StubOpenWebUI) + cls.upstream_url = f"http://127.0.0.1:{cls.upstream.server_address[1]}" + cls.patcher = patch.object(config, "UPSTREAM", cls.upstream_url) + cls.patcher.start() + cls.port_patcher = patch.object(config, "PROXY_PORT", 0) + cls.port_patcher.start() + cls.server = proxy.serve(cls.studio.server_address[1]) + cls.url = f"http://127.0.0.1:{cls.server.server_address[1]}" + + @classmethod + def tearDownClass(cls): + cls.server.shutdown() + cls.studio.shutdown() + cls.upstream.shutdown() + cls.port_patcher.stop() + cls.patcher.stop() + super().tearDownClass() + + def setUp(self): + proxy._identity_cache.clear() + _StubStudio.calls = 0 + + def test_a_signed_in_browser_is_identified(self): + response = httpx.get(f"{self.url}/api/config", headers={"Cookie": VALID_COOKIE}) + self.assertEqual(response.json()["email"], "ada@example.com") + self.assertEqual(response.json()["role"], "admin") + + def test_a_client_cannot_supply_its_own_identity(self): + """In any casing: HTTP header names are case-insensitive. + + Overwriting the headers we forward is not enough on its own. A client + sending `x-studio-email` in another casing would add a second header + rather than replace ours, and the upstream reads whichever comes first. + """ + for name in (config.EMAIL_HEADER, config.EMAIL_HEADER.lower(), config.EMAIL_HEADER.upper()): + response = httpx.get(f"{self.url}/api/config", headers={ + "Cookie": VALID_COOKIE, name: "evil@example.com", + }).json() + self.assertEqual(response["email"], "ada@example.com", name) + self.assertEqual(response["identity_headers_seen"], 1, name) + + def test_a_signed_out_browser_gets_a_way_back_not_the_chat(self): + response = httpx.get(self.url, follow_redirects=False) + self.assertIn("Sign in to Studio", response.text) + self.assertIn("top.location", response.text) # breaks out of the iframe + + def test_one_browsers_session_never_reaches_another(self): + """The proxy must remember nothing between requests. + + This is the bug that made a cookieless request come back as the last + signed-in user: an httpx client keeps a cookie jar, so Studio's sessionid + and Open WebUI's token were stored and replayed for whoever asked next. + """ + signed_in = httpx.get(f"{self.url}/api/config", headers={"Cookie": VALID_COOKIE}) + self.assertEqual(signed_in.json()["email"], "ada@example.com") + + anonymous = httpx.get(f"{self.url}/api/config") + self.assertIn("Sign in to Studio", anonymous.text) + + def test_the_leak_is_what_the_test_above_would_catch(self): + """The control: put the old shared client back, and the leak reappears. + + Without this, the test above passes for any reason at all — including a + broken stub — and would not notice the bug coming back. + """ + shared = httpx.Client(transport=proxy._Handler.transport, follow_redirects=False) + with patch.object(proxy._Handler, "client", lambda self: shared): + httpx.get(f"{self.url}/api/config", headers={"Cookie": VALID_COOKIE}) + anonymous = httpx.get(f"{self.url}/api/config") + self.assertEqual(anonymous.json()["email"], "ada@example.com", + "the shared client no longer leaks — has httpx changed?") + + def test_every_request_gets_a_client_that_remembers_nothing(self): + """The property the fix rests on, asserted without going through HTTP.""" + handler = proxy._Handler.__new__(proxy._Handler) + first, second = handler.client(), handler.client() + self.assertIsNot(first, second) + first.cookies.set("sessionid", "someone-elses") + self.assertEqual(dict(handler.client().cookies), {}) + + def test_the_upstreams_cookies_are_not_kept_either(self): + httpx.get(f"{self.url}/api/config", headers={"Cookie": VALID_COOKIE}) + second = httpx.get(f"{self.url}/api/config", headers={"Cookie": VALID_COOKIE}) + self.assertNotIn("someones-jwt", second.json()["cookie_seen"] or "") + + def test_a_page_load_asks_studio_once_not_once_per_asset(self): + """Open WebUI pulls dozens of assets; each one asking Django would show.""" + for _ in range(5): + httpx.get(f"{self.url}/api/config", headers={"Cookie": VALID_COOKIE}) + self.assertEqual(_StubStudio.calls, 1) + + def test_the_answer_stops_being_used_once_it_is_old(self): + """Otherwise a sign-out would never take effect.""" + with patch.object(proxy, "IDENTITY_TTL", 0.05): + httpx.get(f"{self.url}/api/config", headers={"Cookie": VALID_COOKIE}) + time.sleep(0.1) + httpx.get(f"{self.url}/api/config", headers={"Cookie": VALID_COOKIE}) + self.assertEqual(_StubStudio.calls, 2) + + def test_the_cache_can_be_turned_off(self): + with patch.object(proxy, "IDENTITY_TTL", 0): + for _ in range(3): + httpx.get(f"{self.url}/api/config", headers={"Cookie": VALID_COOKIE}) + self.assertEqual(_StubStudio.calls, 3) + + def test_the_frame_blocking_header_is_removed(self): + response = httpx.get(f"{self.url}/api/config", headers={"Cookie": VALID_COOKIE}) + self.assertNotIn("x-frame-options", response.headers) + + def test_the_embed_stylesheet_comes_from_studio_not_the_upstream(self): + """Open WebUI loads /static/custom.css on every page. + + The proxy answers it with chat/embed.css (which hides the + chat-history sidebar in the iframe) instead of forwarding, so the rule + lives in this repo and survives Open WebUI upgrades. It is a static + asset, so it must not require a signed-in browser. + """ + response = httpx.get(f"{self.url}/static/custom.css") + self.assertEqual(response.status_code, 200) + self.assertIn("text/css", response.headers["Content-Type"]) + self.assertIn("#sidebar", response.text) + self.assertNotIn("email", response.text) # not the upstream's echo + + def test_open_webui_favicon_comes_from_studio_without_auth(self): + response = httpx.get(f"{self.url}/static/favicon-32x32.svg") + expected = (Path(proxy.__file__).resolve().parents[1] / "static" / "logo.svg").read_bytes() + self.assertEqual(response.status_code, 200) + self.assertEqual(response.headers["Content-Type"], "image/svg+xml") + self.assertEqual(response.content, expected) + + @contextmanager + def _tunnel_upstream(self, reply: bytes): + """A fresh stub upstream that answers the upgrade with ``reply``.""" + _StubWebSocketUpstream.reply = reply + _StubWebSocketUpstream.saw_identity = False + server = _serve(_StubWebSocketUpstream) + with patch.object(config, "UPSTREAM", + f"http://127.0.0.1:{server.server_address[1]}"): + yield server + server.shutdown() + + def test_a_refused_upgrade_comes_back_as_an_http_error_not_a_websocket(self): + """A 401 from the upstream must reach the browser as a 401. + + Before the fix the raw error bytes were piped as if they were + WebSocket frames, and the browser failed opaquely. + """ + reply = ( + b"HTTP/1.1 401 Unauthorized\r\n" + b"Content-Type: text/plain\r\n" + b"Content-Length: 13\r\n" + b"\r\n" + b"unauthorized\n" + ) + with self._tunnel_upstream(reply): + response = httpx.get( + f"{self.url}/ws", + headers={ + "Cookie": VALID_COOKIE, + "Connection": "Upgrade", + "Upgrade": "websocket", + "Sec-WebSocket-Version": "13", + "Sec-WebSocket-Key": "dGhlIHNhbXBsZSBub25jZQ==", + }, + timeout=10, + ) + self.assertEqual(response.status_code, 401) + self.assertEqual(response.text, "unauthorized\n") + self.assertTrue(_StubWebSocketUpstream.saw_identity) + + def test_a_successful_upgrade_is_piped_both_ways_without_losing_bytes(self): + """A 101 is tunnelled, and bytes sent with the 101 are not lost. + + The stub answers the handshake and, in the same packet, sends the + first "frame" bytes; the tunnel must deliver them, and must carry + bytes back the other way too. + """ + first = b"hello-from-upstream" + reply = ( + b"HTTP/1.1 101 Switching Protocols\r\n" + b"Upgrade: websocket\r\n" + b"Connection: Upgrade\r\n" + b"\r\n" + ) + first + with self._tunnel_upstream(reply), socket.create_connection( + ("127.0.0.1", self.server.server_address[1]), timeout=10 + ) as client: + client.sendall( + b"GET /ws HTTP/1.1\r\n" + b"Host: 127.0.0.1\r\n" + b"Connection: Upgrade\r\n" + b"Upgrade: websocket\r\n" + b"Sec-WebSocket-Version: 13\r\n" + b"Sec-WebSocket-Key: dGhlIHNhbXBsZSBub25jZQ==\r\n" + b"Cookie: " + VALID_COOKIE.encode() + b"\r\n" + b"\r\n" + ) + # The 101 header, then the first bytes the stub sent with it. + buffer = b"" + while not buffer.endswith(first): + chunk = client.recv(65536) + self.assertTrue(chunk, "the tunnel closed before the first bytes") + buffer += chunk + self.assertIn(b"101 Switching Protocols", buffer) + self.assertTrue(buffer.endswith(first)) + + # Client to upstream: the stub echoes it straight back. + client.sendall(b"ping-from-client") + buffer = b"" + while b"ping-from-client" not in buffer: + chunk = client.recv(65536) + self.assertTrue(chunk, "the tunnel closed before the echo") + buffer += chunk + self.assertTrue(_StubWebSocketUpstream.saw_identity) + + +class ReadinessTests(SimpleTestCase): + """wait_until_ready must not be fooled by *another* Open WebUI on the port. + + The health URL is fixed, so a 200 can come from an instance that was + already there when ours died on the bind. + """ + + def setUp(self): + self.upstream = _serve(_StubOpenWebUI) + self.patcher = patch.object(config, "UPSTREAM", + f"http://127.0.0.1:{self.upstream.server_address[1]}") + self.patcher.start() + + def tearDown(self): + self.patcher.stop() + self.upstream.shutdown() + super().tearDown() + + def test_a_200_from_a_dead_process_is_not_readiness(self): + process = MagicMock() + process.poll.return_value = 0 # ours died on the bind + self.assertFalse(proxy.wait_until_ready(process, timeout=5)) + + def test_a_200_from_a_live_process_is_readiness(self): + process = MagicMock() + process.poll.return_value = None + self.assertTrue(proxy.wait_until_ready(process, timeout=5)) + + +class PidFileTests(SimpleTestCase): + """The pid file must never trample or orphan another run's Open WebUI.""" + + def setUp(self): + self._tmp = tempfile.TemporaryDirectory() + self.tmp = Path(self._tmp.name) + self.file_patcher = patch.object(proxy, "pid_file", lambda: self.tmp / "open-webui.pid") + self.file_patcher.start() + self._process = proxy._process + proxy._process = None + + def tearDown(self): + proxy._process = self._process + self.file_patcher.stop() + self._tmp.cleanup() + super().tearDown() + + def test_stop_does_not_delete_a_file_recording_someone_elses_pid(self): + file = self.tmp / "open-webui.pid" + file.write_text("424242") + process = MagicMock() + process.pid = 1234 + process.poll.return_value = 0 # already dead: nothing to terminate + proxy._process = process + proxy.stop_open_webui() + self.assertEqual(file.read_text().strip(), "424242") + + def test_stop_deletes_a_file_recording_our_pid(self): + file = self.tmp / "open-webui.pid" + process = MagicMock() + process.pid = 1234 + process.poll.return_value = 0 + proxy._process = process + file.write_text("1234") + proxy.stop_open_webui() + self.assertFalse(file.exists()) + + def test_spawn_does_not_overwrite_a_file_recording_a_live_pid(self): + file = self.tmp / "open-webui.pid" + file.write_text("424242") + fake = MagicMock() + fake.pid = 1234 + with patch.object(proxy.subprocess, "Popen", return_value=fake), \ + patch.object(proxy, "log_path", lambda: self.tmp / "server.log"), \ + patch("infra.minimal_config._pid_alive", return_value=True): + proxy._spawn(self.tmp) + self.assertEqual(file.read_text().strip(), "424242") + proxy._process = None + + def test_spawn_overwrites_a_file_recording_a_dead_pid(self): + file = self.tmp / "open-webui.pid" + file.write_text("424242") + fake = MagicMock() + fake.pid = 1234 + with patch.object(proxy.subprocess, "Popen", return_value=fake), \ + patch.object(proxy, "log_path", lambda: self.tmp / "server.log"), \ + patch("infra.minimal_config._pid_alive", return_value=False): + proxy._spawn(self.tmp) + self.assertEqual(file.read_text().strip(), "1234") + proxy._process = None + + def test_stop_stale_leaves_a_live_parented_process_alone(self): + """A recorded PID with a live parent belongs to another running run.""" + file = self.tmp / "open-webui.pid" + file.write_text(str(os.getpid())) # alive, parented by this test process + with patch("infra.minimal_config._pid_alive", return_value=True), \ + patch("infra.minimal_config._parent_pid", return_value=os.getpid()), \ + patch("infra.minimal_config._orphaned", return_value=False), \ + patch.object(proxy, "_terminate") as terminate: + proxy._stop_stale() + terminate.assert_not_called() + self.assertTrue(file.exists()) + + +class OtlpEnvTests(SimpleTestCase): + """When OTLP is enabled, Open WebUI is pointed at Studio's listener.""" + + def setUp(self): + self._tmp = tempfile.TemporaryDirectory() + self.tmp = Path(self._tmp.name) + self._process = proxy._process + proxy._process = None + self._pid_patcher = patch.object(proxy, "pid_file", lambda: self.tmp / "open-webui.pid") + self._pid_patcher.start() + + def tearDown(self): + self._pid_patcher.stop() + proxy._process = self._process + self._tmp.cleanup() + super().tearDown() + + def _spawn_env(self, endpoint=""): + """Run _spawn with OTLP enabled and return the env it was given.""" + fake = MagicMock() + fake.pid = 1234 + with patch.object(proxy.subprocess, "Popen", return_value=fake) as popen, \ + patch.object(proxy, "log_path", lambda: self.tmp / "server.log"), \ + patch("infra.minimal_config._pid_alive", return_value=False), \ + patch.object(config, "OTLP_ENABLED", True), \ + patch.object(config, "OTLP_ENDPOINT", endpoint), \ + patch.object(config, "OTLP_SERVICE_NAME", "open-webui"): + proxy._spawn(self.tmp, 8000) + proxy._process = None + return popen.call_args.kwargs["env"] + + def test_disabled_by_default_adds_no_otel_vars(self): + fake = MagicMock() + fake.pid = 1234 + with patch.object(proxy.subprocess, "Popen", return_value=fake) as popen, \ + patch.object(proxy, "log_path", lambda: self.tmp / "server.log"), \ + patch("infra.minimal_config._pid_alive", return_value=False), \ + patch.object(config, "OTLP_ENABLED", False): + proxy._spawn(self.tmp, 8000) + proxy._process = None + env = popen.call_args.kwargs["env"] + self.assertNotIn("ENABLE_OTEL", env) + self.assertNotIn("OTEL_EXPORTER_OTLP_ENDPOINT", env) + + def test_enabled_points_at_studios_listener(self): + env = self._spawn_env() + self.assertEqual(env["ENABLE_OTEL"], "true") + self.assertEqual(env["ENABLE_OTEL_TRACES"], "true") + self.assertEqual(env["OTEL_OTLP_SPAN_EXPORTER"], "http") + # Base URL only — Open WebUI's exporter appends /v1/traces itself. + self.assertEqual(env["OTEL_EXPORTER_OTLP_ENDPOINT"], "http://127.0.0.1:8000") + self.assertEqual(env["OTEL_SERVICE_NAME"], "open-webui") + + def test_explicit_endpoint_wins(self): + env = self._spawn_env(endpoint="http://collector:4318") + self.assertEqual(env["OTEL_EXPORTER_OTLP_ENDPOINT"], "http://collector:4318") diff --git a/chat/tests/test_sync.py b/chat/tests/test_sync.py new file mode 100644 index 00000000..0ab86681 --- /dev/null +++ b/chat/tests/test_sync.py @@ -0,0 +1,228 @@ +"""Connections change in Studio, chat follows — without blocking the save. + +Run: + SIMPLEAUDIT_LOCAL_SQLITE=1 uv run manage.py test chat.tests.test_sync +""" +import time +from unittest import mock +from unittest.mock import patch + +from django.test import TestCase, TransactionTestCase + +from chat import sync +from chat.api import ChatAPIError +from chat.sync import NoAdminError +from infra.tests.factories import ( + ModelConnectionFactory, + ProjectFactory, + RegisteredModelFactory, + UserFactory, +) + + +class ConnectionsToPushTests(TestCase): + def setUp(self): + self.project = ProjectFactory() + + def test_only_enabled_connections_with_a_base_url(self): + ModelConnectionFactory(project=self.project, name="on", base_url="https://a.example/v1") + ModelConnectionFactory(project=self.project, name="off", base_url="https://b.example/v1", + enabled=False) + ModelConnectionFactory(project=self.project, name="blank", base_url="") + self.assertEqual([c["name"] for c in sync.connections_to_push()], ["on"]) + + def test_registered_models_narrow_the_connection(self): + conn = ModelConnectionFactory(project=self.project, base_url="https://a.example/v1") + RegisteredModelFactory(connection=conn, project=self.project, model_id="gpt-4o") + RegisteredModelFactory(connection=conn, project=self.project, model_id="gpt-4o-mini") + RegisteredModelFactory(connection=conn, project=self.project, model_id="gone", enabled=False) + self.assertEqual(sync.connections_to_push()[0]["model_ids"], ["gpt-4o", "gpt-4o-mini"]) + + def test_a_connection_without_registered_models_is_not_narrowed(self): + ModelConnectionFactory(project=self.project, base_url="https://a.example/v1") + self.assertEqual(sync.connections_to_push()[0]["model_ids"], []) + + +class SchedulePushTests(TestCase): + """The debounce: many changes, one push, off the caller's thread.""" + + def setUp(self): + self.pushes = [] + patcher = patch.object(sync, "push_now", lambda: self.pushes.append(1) or {"pushed": 1, "kept": 0}) + patcher.start() + self.addCleanup(patcher.stop) + delay = patch.object(sync, "DELAY_SECONDS", 0.05) + delay.start() + self.addCleanup(delay.stop) + + def _settle(self): + time.sleep(0.3) + + @patch("chat.config.ENABLED", True) + def test_a_burst_of_changes_is_one_push(self): + for _ in range(5): + sync.schedule_push("test") + self._settle() + self.assertEqual(len(self.pushes), 1) + + @patch("chat.config.ENABLED", False) + def test_nothing_is_pushed_while_chat_is_disabled(self): + sync.schedule_push("test") + self._settle() + self.assertEqual(self.pushes, []) + + @patch("chat.config.ENABLED", True) + def test_a_chat_that_is_not_answering_does_not_raise(self): + with patch.object(sync, "push_now", side_effect=ChatAPIError("not running")): + sync.schedule_push("test") + self._settle() # the failure is logged, the caller never sees it + + +class RunTests(TestCase): + """_run is the background worker: it must never raise, and it must tell + a misconfiguration (no superuser) apart from a transient skip.""" + + def test_no_superuser_is_logged_loudly_with_an_actionable_message(self): + with patch.object( + sync, "push_now", + side_effect=NoAdminError("No superuser to act as; Open WebUI's provider config needs an admin."), + ), self.assertLogs("chat.sync", level="WARNING") as captured: + sync._run("test") + self.assertTrue( + any("no Studio superuser" in line for line in captured.output), + captured.output, + ) + self.assertTrue( + any("manage.py sync_chat_models" in line for line in captured.output), + captured.output, + ) + + def test_a_transient_failure_is_a_quiet_skip_not_a_no_superuser_warning(self): + with patch.object(sync, "push_now", side_effect=ChatAPIError("not running")), \ + self.assertLogs("chat.sync", level="INFO") as captured: + sync._run("test") + self.assertTrue( + any("Chat model sync skipped" in line and "not running" in line for line in captured.output), + captured.output, + ) + self.assertFalse( + any("no Studio superuser" in line for line in captured.output), + captured.output, + ) + + +class ReconcileModelIdsTests(TestCase): + """push_now warns when a pinned Studio model id is not one Open WebUI + registers — the case where the ?models= pin would silently no-op.""" + + def _push(self, registered, model_ids): + api = mock.Mock() + api.push_connections.return_value = {"pushed": 1, "kept": 0} + api.list_models.return_value = registered + with patch.object(sync, "ChatAPI") as chat_api, \ + patch.object(sync, "_admin", return_value=object()): + chat_api.as_user.return_value = api + sync.push_now([ + {"id": 1, "name": "OpenAI", "base_url": "https://api.openai.com/v1", + "api_key": "sk-x", "enabled": True, "model_ids": model_ids}, + ]) + return api + + def test_a_missing_model_id_is_logged_loudly(self): + with self.assertLogs("chat.sync", level="WARNING") as captured: + api = self._push(["gpt-4o"], ["gpt-4o", "studio-internal-id"]) + self.assertTrue( + any("will not work for connection 'OpenAI'" in line + and "studio-internal-id" in line + for line in captured.output), + captured.output, + ) + api.list_models.assert_called_once() + + def test_all_ids_registered_logs_no_warning(self): + import logging + + records = [] + + class _Capture(logging.Handler): + def emit(self, record): + records.append(record.getMessage()) + + handler = _Capture(level=logging.WARNING) + with mock.patch.object(logging.getLogger("chat.sync"), "handlers", [handler]): + self._push(["gpt-4o", "gpt-4o-mini"], ["gpt-4o", "gpt-4o-mini"]) + self.assertFalse( + any("will not work" in line for line in records), + records, + ) + + def test_a_connection_without_model_ids_is_not_checked(self): + api = self._push(["gpt-4o"], []) + api.list_models.assert_not_called() + + def test_a_models_endpoint_failure_skips_reconciliation(self): + api = mock.Mock() + api.push_connections.return_value = {"pushed": 1, "kept": 0} + api.list_models.side_effect = ChatAPIError("not ready") + with patch.object(sync, "ChatAPI") as chat_api, \ + patch.object(sync, "_admin", return_value=object()): + chat_api.as_user.return_value = api + sync.push_now([ + {"id": 1, "name": "OpenAI", "base_url": "https://api.openai.com/v1", + "api_key": "sk-x", "enabled": True, "model_ids": ["gpt-4o"]}, + ]) + # The push itself still reports success. + self.assertEqual(api.push_connections.call_count, 1) + + +class SignalTests(TransactionTestCase): + """on_commit means the push waits for the transaction, and skips a rollback. + + Importing chat.signals is what connects the receivers (that is what + ChatConfig.ready does when chat is on), so the test does it explicitly rather + than depending on the environment it runs in. + """ + + def setUp(self): + import chat.signals # noqa: F401 (connects the receivers) + + self.project = ProjectFactory() + self.scheduled = [] + patcher = patch.object(sync, "schedule_push", lambda reason="": self.scheduled.append(reason)) + patcher.start() + self.addCleanup(patcher.stop) + + def test_saving_a_connection_schedules_a_push(self): + conn = ModelConnectionFactory(project=self.project, base_url="https://a.example/v1") + self.assertTrue(any(str(conn.pk) in reason for reason in self.scheduled)) + + def test_deleting_a_connection_schedules_a_push(self): + conn = ModelConnectionFactory(project=self.project, base_url="https://a.example/v1") + self.scheduled.clear() + conn.delete() + self.assertEqual(len(self.scheduled), 1) + + def test_registering_a_model_schedules_a_push(self): + conn = ModelConnectionFactory(project=self.project, base_url="https://a.example/v1") + self.scheduled.clear() + RegisteredModelFactory(connection=conn, project=self.project, model_id="gpt-4o") + self.assertEqual(len(self.scheduled), 1) + + def test_a_rolled_back_change_pushes_nothing(self): + from django.db import transaction + + try: + with transaction.atomic(): + ModelConnectionFactory(project=self.project, base_url="https://a.example/v1") + raise RuntimeError("rolled back") + except RuntimeError: + pass + self.assertEqual(self.scheduled, []) + + +class PushNowTests(TestCase): + def test_without_a_superuser_it_says_so(self): + UserFactory(username="ordinary", is_superuser=False) + with self.assertRaises(ChatAPIError) as caught: + sync.push_now() + self.assertIn("No superuser", str(caught.exception)) diff --git a/chat/urls.py b/chat/urls.py new file mode 100644 index 00000000..286337d4 --- /dev/null +++ b/chat/urls.py @@ -0,0 +1,14 @@ +"""Chat URLs, mounted at /chat/ by config/urls.py. + +Both views 404 while the module is disabled, so they are safe to include +unconditionally. +""" +from django.urls import path + +from chat.views import ChatView, authz, chat_with + +urlpatterns = [ + path("", ChatView.as_view(), name="chat"), + path("with//", chat_with, name="chat_with"), + path("authz", authz, name="chat_authz"), +] diff --git a/chat/views.py b/chat/views.py new file mode 100644 index 00000000..22877a18 --- /dev/null +++ b/chat/views.py @@ -0,0 +1,206 @@ +"""The two things Studio serves for chat: the page, and who the browser is. + +Everything about *why* it works this way is in chat/config.py. +""" +from __future__ import annotations + +from urllib.parse import urlencode + +from django.http import Http404, HttpResponse +from django.shortcuts import redirect +from django.urls import reverse +from django.views.generic import TemplateView + +# The module, not the names: tests and runtime both read ENABLED/PUBLIC_URL as +# they are now, not as they were at import time. +from chat import config + + +def authz(request): + """Forward-auth endpoint: who is this browser? + + 200 with the trusted headers when signed in, 401 otherwise. The proxy copies + the headers onto the upstream request and turns a 401 into a redirect to + Studio's login page. + """ + if not config.ENABLED: + raise Http404 + user = request.user + if not user.is_authenticated: + return HttpResponse(status=401) + response = HttpResponse(status=200) + for header, value in config.identity(user).items(): + response[header] = value + return response + + +def chat_with(request, connection_id, model_id): + """Hand off from /connections/ to the chat, pinned to one model. + + The icon links here (not straight to /chat/?model=) so the model is + validated server-side and carried in the session, where ChatView consumes + it once. A hand-typed ?model= on /chat/ is ignored — only a model the user + can actually see, chosen through this view, can pin the chat. + + The model is identified by its connection plus its bare model id, because + the same model id can be registered under more than one connection. The + session stores the prefixed id (``.``) that Open + WebUI expects once Studio pushes a ``prefix_id`` per connection. + """ + if not config.ENABLED: + raise Http404 + if not request.user.is_authenticated: + return redirect(f"/login/?next={request.path}") + model_id = (model_id or "").strip() + if model_id and _user_can_see_model(request, connection_id, model_id): + request.session["chat_pinned_model"] = f"{connection_id}.{model_id}" + return redirect(reverse("chat")) + + +def _user_can_see_model(request, connection_id, model_id) -> bool: + """True if the user has a visible, enabled connection serving model_id. + + The connection must be one the user can see in this workspace, enabled, and + actually serving the bare model id — so a hand-typed id for a connection the + user cannot see is rejected even if the bare id exists elsewhere. + """ + from model_registry.services import visible_connections_for + + project = getattr(request, "project", None) + if project is None: + return False + conn = next( + (c for c in visible_connections_for(request.user, project) if c.id == connection_id), + None, + ) + if conn is None or not conn.enabled or not (conn.base_url or "").strip(): + return False + return conn.models.filter(enabled=True, model_id=model_id).exists() + + +class ChatView(TemplateView): + """The Studio page that embeds Open WebUI.""" + + template_name = "chat/chat.html" + + def get(self, request, *args, **kwargs): + if not config.ENABLED: + raise Http404 + if not request.user.is_authenticated: + return redirect(f"/login/?next={request.path}") + return super().get(request, *args, **kwargs) + + def get_context_data(self, **kwargs): + chat_base = config.public_url(self.request) + # Which model(s) to pin. A model chosen through the /connections chat + # icon is stashed in the session by chat_with and consumed here exactly + # once (pop) — so it survives the redirect but not a refresh, and a + # hand-typed ?model= can never pin the chat. Otherwise the user's saved + # preference, or the first model they can see. + pinned = self.request.session.pop("chat_pinned_model", None) or self._resolve_default_model() + # Shape the embedded chat through URL params, which Open WebUI reads on + # load: + # ?models= pin to one or more models, comma-separated (the + # in-frame picker is hidden, so this is the only way + # to choose them). The top-bar picker rebuilds this. + # ?temporary-chat start in temporary mode, so nothing is saved to the + # chat history. The embed is a throwaway surface. + # The New Chat button is hidden, so a fresh chat only ever starts from a + # full page load — which re-reads these params — so the param is enough. + params = {} + if pinned: + params["models"] = pinned + params["temporary-chat"] = "true" + chat_url = f"{chat_base}?{urlencode(params)}" + groups = self._chat_model_groups() + # The full-page "no models" state is only for a user with no project at + # all (request.project is None) — that is the case where the chat is + # genuinely unusable and there is no workspace to point at. A user who + # has a project but no models yet still gets the normal page: the + # picker shows its own "No models yet" hint and the iframe loads. + has_project = getattr(self.request, "project", None) is not None + context = { + "chat_url": chat_url, + "chat_model_data": { + "base": chat_base, + "groups": groups, + "defaults": [m.strip() for m in pinned.split(",") if m.strip()], + }, + "has_models": has_project, + } + if not has_project: + context["no_models_message"] = ( + "No models are available for this workspace yet. " + "Register a connection to start chatting." + ) + return super().get_context_data(**context, **kwargs) + + def _visible_models(self): + """The models the user can chat with, as ``{"id", "name", "has_key"}``. + + Same visibility rule as the experiments page: this workspace's own + connections plus any shared into it, each enabled and reachable. Each + id is the prefixed id (``.``) — that is what + Open WebUI's ?model=/ ?models= params expect once Studio pushes a + ``prefix_id`` per connection, and it is what disambiguates the same + model id under two connections. + """ + from chat.api import chat_model_prefix + from model_registry.services import visible_connections_for + + project = getattr(self.request, "project", None) + if project is None: + return [] + models = [] + for conn in visible_connections_for(self.request.user, project): + if not conn.enabled or not (conn.base_url or "").strip(): + continue + prefix = chat_model_prefix(conn) + models.extend( + {"id": f"{prefix}.{m.model_id}", "name": m.display_name, "has_key": m.has_key} + for m in conn.models.filter(enabled=True) + ) + return models + + def _chat_model_groups(self): + """The models the top-bar picker offers, grouped by connection.""" + from chat.api import chat_model_prefix + from model_registry.services import visible_connections_for + + project = getattr(self.request, "project", None) + if project is None: + return [] + groups = [] + for conn in visible_connections_for(self.request.user, project): + if not conn.enabled or not (conn.base_url or "").strip(): + continue + prefix = chat_model_prefix(conn) + models = [ + {"id": f"{prefix}.{m.model_id}", "name": m.display_name, "has_key": m.has_key} + for m in conn.models.filter(enabled=True) + ] + if models: + groups.append({"connection": conn.name, "models": models}) + return groups + + def _resolve_default_model(self): + """The model(s) to pin when no model was chosen through the handoff. + + The user's saved preference wins if it still names models they can see + (a connection may have been deleted or a model removed since) — the + saved value is the comma-separated selection, and any ids that are no + longer available are dropped. Otherwise the first available model is + pinned, so the chat opens on something usable rather than an empty + picker. If the user has no visible models at all, nothing is pinned — + Open WebUI shows its own default. + """ + available = self._visible_models() + if not available: + return "" + available_ids = {m["id"] for m in available} + saved = (self.request.user.preferences or {}).get("chat_model") + if saved: + saved_ids = [i for i in (s.strip() for s in saved.split(",")) if i in available_ids] + if saved_ids: + return ",".join(saved_ids) + return available[0]["id"] diff --git a/config/settings.py b/config/settings.py index 8440b419..44c53245 100644 --- a/config/settings.py +++ b/config/settings.py @@ -120,6 +120,10 @@ def _csrf_trusted_origins() -> list[str]: "judges", "audits", "infra", + # Optional module: its URLs 404 and nothing runs unless SIMPLEAUDIT_CHAT is + # set (chat/config.py). Installed either way so its templates, management + # commands and tests resolve. + "chat", ] MIDDLEWARE = [ @@ -296,6 +300,26 @@ def _csrf_trusted_origins() -> list[str]: SESSION_COOKIE_SECURE = True CSRF_COOKIE_SECURE = True +# Production security headers (per Django's production checklist). The app runs +# behind a TLS-terminating proxy (see SECURE_PROXY_SSL_HEADER above), so these +# are safe to enable. They are gated off for local dev (http://localhost), the +# local minimal/demo bundle, DEMO_MODE (which sets its own cookie policy), and +# the test suite (which runs over http://localhost with DEBUG=false) so they +# never break a non-HTTPS or cross-site-embedded setup. +_TESTING = ( + SECRET_KEY in {"test-secret-key-not-change-me", "ci-secret-key"} + or bool(os.environ.get("PYTEST_CURRENT_TEST")) + or env_bool("SIMPLEAUDIT_TESTING", False) +) +if not (DEBUG or MINIMAL_CONFIG or DEMO_MODE or _TESTING): + SECURE_SSL_REDIRECT = True + SECURE_HSTS_SECONDS = 60 * 60 * 24 * 365 # 1 year + SECURE_HSTS_INCLUDE_SUBDOMAINS = True + SECURE_HSTS_PRELOAD = True + SESSION_COOKIE_SECURE = True + CSRF_COOKIE_SECURE = True + SECURE_CONTENT_TYPE_NOSNIFF = True + # Operational settings used by health checks and bootstrap commands. if MINIMAL_CONFIG: # In local demo mode the embedded Hatchet client provides its own connection diff --git a/config/urls.py b/config/urls.py index 74c0a185..c0714a6f 100644 --- a/config/urls.py +++ b/config/urls.py @@ -31,6 +31,8 @@ MonitorDetailView, MonitorsView, NewExperimentView, + OTLPCredentialCreateView, + OTLPCredentialRotateView, ProfileView, RegisterView, RunArchiveView, @@ -59,6 +61,7 @@ logout_view, ) from judges.views import JudgeDetailView, JudgePreviewView, JudgesView +from model_registry import otlp_config, otlp_views # --- Static file serving --------------------------------------------------- # For the canonical Docker Compose self-hosted deployment Django serves its @@ -140,6 +143,10 @@ def home_view(request, *args, **kwargs): path("api/", include("scenarios.urls")), path("api/", include("model_registry.urls")), path("api/", include("audits.urls")), + # Shared OTLP ingestion endpoint (machine-to-machine, Basic/****** + # Only wired while the OTLP listener is on (SIMPLEAUDIT_OTLP); otherwise + # these paths 404 so a deployment that doesn't want the listener exposes + # no OTLP surface at all. See model_registry/otlp_config.py. # Public landing page (indexable, no auth) — the site's SEO surface # "/" is the dashboard when signed in and the public landing page otherwise. path("", home_view, name="dashboard"), @@ -147,7 +154,7 @@ def home_view(request, *args, **kwargs): # UI (server-rendered CBVs) path("login/", LoginView.as_view(), name="login"), # Local one-liner demo only (404 unless MINIMAL_CONFIG): the CLI opens this - # in the default browser to land the user signed-in on the dashboard. + # in the default browser with a single-use ?token=... to sign the user in. path("auto-login/", auto_login_view, name="auto_login"), path("register/", RegisterView.as_view(), name="register"), path("logout/", logout_view, name="logout"), @@ -183,6 +190,8 @@ def home_view(request, *args, **kwargs): path("connections/discover/", DiscoverModelsView.as_view(), name="models_discover"), path("connections/check/", ConnectionCheckView.as_view(), name="connection_check"), path("connections//delete/", ConnectionDeleteView.as_view(), name="connection_delete"), + path("connections/otlp-credential/", OTLPCredentialCreateView.as_view(), name="otlp_credential_create"), + path("connections/otlp-credential/rotate/", OTLPCredentialRotateView.as_view(), name="otlp_credential_rotate"), path("judges/", JudgesView.as_view(), name="judges"), path("judges/new/", JudgeDetailView.as_view(), name="judge_new"), path("judges/preview/", JudgePreviewView.as_view(), name="judge_preview"), @@ -202,3 +211,20 @@ def home_view(request, *args, **kwargs): path("runs/export.csv", RunsExportView.as_view(), name="runs_export"), path("me/preferences/", PreferenceView.as_view(), name="preferences"), ] + +# Optional Open WebUI module (the `chat` app). The chat UI itself lives on its +# own origin; these routes are the iframe page and the forward-auth endpoint its +# proxy calls. Both 404 unless SIMPLEAUDIT_CHAT is set. See chat/config.py. +urlpatterns += [path("chat/", include("chat.urls"))] + +# OTLP listener (machine-to-machine span ingestion + credential management). +# Wired only while SIMPLEAUDIT_OTLP is on (the default); otherwise the +# /otlp/* and /api/otlp/* paths 404. See model_registry/otlp_config.py. +if otlp_config.ENABLED: + urlpatterns += [ + path("otlp/v1/traces", otlp_views.otlp_traces, name="otlp-traces"), + path("api/otlp/credentials/", otlp_views.list_credentials, name="otlp-credentials-list"), + path("api/otlp/credentials/create/", otlp_views.create_credential, name="otlp-credentials-create"), + path("api/otlp/credentials//revoke/", otlp_views.revoke_credential, name="otlp-credentials-revoke"), + ] + diff --git a/conftest.py b/conftest.py new file mode 100644 index 00000000..b5eaae74 --- /dev/null +++ b/conftest.py @@ -0,0 +1,31 @@ +"""Pytest bootstrap for the SimpleAuditStudio Django test suite. + +The existing tests are written as ``django.test.TestCase`` classes and are run +unchanged by pytest-django. The environment the settings module needs is set in +``pytest.ini`` (parsed before this file). This hook mirrors each test's Django +``@tag("...")`` values onto the pytest item as markers, so ``pytest -m "not +slow"`` works the same way ``manage.py test --exclude-tag slow`` does. +""" + + +def pytest_collection_modifyitems(items): + """Copy Django ``@tag`` values onto pytest items as markers. + + Django's ``tag()`` decorator stores tag names in ``unittest``'s + ``_testcase_tags`` attribute on the class (and, when applied to a method, + in the function's ``_testcase_tags``). We read that and attach a matching + pytest marker to every test item, so ``-m`` selection and ``--testmon`` + both see the same grouping the Django runner uses. + """ + for item in items: + tags = set() + # Class-level tags (e.g. @tag("slow") on a TestCase subclass). + cls = getattr(item, "cls", None) + if cls is not None: + tags.update(getattr(cls, "_testcase_tags", ())) + # Method-level tags. + func = getattr(item, "function", None) + if func is not None: + tags.update(getattr(func, "_testcase_tags", ())) + for tag in tags: + item.add_marker(tag) diff --git a/deploy/compose/Caddyfile.chat b/deploy/compose/Caddyfile.chat new file mode 100644 index 00000000..1a2078cd --- /dev/null +++ b/deploy/compose/Caddyfile.chat @@ -0,0 +1,53 @@ +# Forward-auth proxy in front of Open WebUI (docker mode of the optional chat +# module; see infra/chat.py). Caddy asks Studio who the browser is and injects +# the answer as the trusted headers Open WebUI reads. +# +# Open WebUI has no published port in docker-compose.yml, so this proxy is the +# only way to reach it. Keep it that way: a client that can talk to it directly +# can set X-Studio-Role: admin and take over the instance. +{ + auto_https off + admin off +} + +:{$SIMPLEAUDIT_CHAT_PROXY_PORT:8801} { + # Never let a client supply its own identity. + request_header -X-Studio-Email + request_header -X-Studio-Name + request_header -X-Studio-Role + + forward_auth {$STUDIO_UPSTREAM:web:8000} { + uri /chat/authz + copy_headers X-Studio-Email X-Studio-Name X-Studio-Role + + # Signed out: send the browser to Studio rather than showing a bare 401. + @signedout status 401 + handle_response @signedout { + redir {$STUDIO_URL:http://localhost:8000}/chat/ 302 + } + } + + # Studio embeds this origin in an iframe. + header -X-Frame-Options + + # Open WebUI requests favicon files from absolute /static paths. Serve the + # Studio logo mounted at /etc/caddy/branding so Docker and embedded mode share + # the same browser tab branding. + @studio_favicon path /favicon* /static/favicon* + handle @studio_favicon { + root * /etc/caddy/branding + rewrite * /logo.svg + file_server + } + + # Studio's embed stylesheet: Open WebUI loads /static/custom.css on every + # page. Serve it from the repo (mounted read-only) so the iframe renders + # without the chat-history sidebar and the rule survives Open WebUI upgrades. + # Same file the no-Docker proxy serves, so both modes share one source. + handle /static/custom.css { + root * /etc/caddy/embed + file_server + } + + reverse_proxy open-webui:{$SIMPLEAUDIT_CHAT_UPSTREAM_PORT:8080} +} diff --git a/deploy/compose/Dockerfile b/deploy/compose/Dockerfile index 1b742f5a..944e5a83 100644 --- a/deploy/compose/Dockerfile +++ b/deploy/compose/Dockerfile @@ -23,18 +23,22 @@ ENV PYTHONDONTWRITEBYTECODE=1 \ WORKDIR /app -# curl: healthchecks +# curl: healthchecks; git: clone the core (local path dependency) RUN apt-get update \ && apt-get install -y --no-install-recommends \ curl \ + git \ && rm -rf /var/lib/apt/lists/* # pyproject.toml is the single source of truth; uv.lock pins exact versions. # The compose stack runs against external Postgres (postgres extra) and serves # via gunicorn (server extra). uv sync creates /app/.venv; the PATH update # keeps the `python`/`gunicorn` entrypoints working. +# The core (simpleaudit) is a local path dependency (../SimpleAudit) so the +# studio can use the tracing layer not yet in the published wheel. COPY pyproject.toml uv.lock README.md ./ -RUN pip install uv \ +RUN git clone --depth 1 https://github.com/kelkalot/simpleaudit.git /SimpleAudit \ + && pip install uv \ && uv sync --frozen --no-install-project --no-dev --extra postgres --extra server ENV PATH="/app/.venv/bin:$PATH" diff --git a/docker-compose.yml b/docker-compose.yml index 5a2157f3..d34ac9d2 100644 --- a/docker-compose.yml +++ b/docker-compose.yml @@ -160,8 +160,74 @@ services: interval: 30s timeout: 5s retries: 5 + # Matches the worker. The whole command is idempotent (migrate, bootstrap + # and seed all skip what already exists), so a restart is safe and a + # gunicorn crash or a Postgres blip no longer leaves the UI down for good. + # A bad .env still fails every attempt, but the error repeats in the logs + # instead of being printed once into a container nobody looks at. + restart: unless-stopped + + # Optional chat module. .env carries both switches (SIMPLEAUDIT_CHAT=docker so + # the web service serves /chat/, COMPOSE_PROFILES=chat so these two containers + # are part of a plain `docker compose up -d`); comment them out to leave chat + # out entirely. See infra/chat.py. + # + # Open WebUI deliberately publishes NO port: chat-proxy is the only route to + # it, because it trusts the identity headers on any request it receives. + open-webui: + profiles: ["chat"] + image: ${OPEN_WEBUI_IMAGE:-ghcr.io/open-webui/open-webui:main} + environment: + WEBUI_AUTH_TRUSTED_EMAIL_HEADER: X-Studio-Email + WEBUI_AUTH_TRUSTED_NAME_HEADER: X-Studio-Name + WEBUI_AUTH_TRUSTED_ROLE_HEADER: X-Studio-Role + ENABLE_SIGNUP: "false" + # Nothing in this stack serves Ollama; Open WebUI polls it otherwise. + ENABLE_OLLAMA_API: "false" + WEBUI_URL: ${SIMPLEAUDIT_CHAT_URL:-http://localhost:8801} + # The image entrypoint is start.sh (no `open-webui` CLI on PATH); it reads + # HOST/PORT and launches uvicorn itself. + HOST: "0.0.0.0" + PORT: ${SIMPLEAUDIT_CHAT_UPSTREAM_PORT:-8080} + # Optional: export Open WebUI's spans to Studio's OTLP listener. Off by + # default — set SIMPLEAUDIT_CHAT_OTLP=true to enable. Open WebUI picks the + # HTTP exporter from OTEL_OTLP_SPAN_EXPORTER and appends /v1/traces to the + # endpoint, so the endpoint is the base URL (Studio's web origin). Exports + # unauthenticated by default (matches an enabled "none" credential); for a + # basic/bearer credential set OTEL_BASIC_AUTH_USERNAME / OTEL_BASIC_AUTH_PASSWORD. + ENABLE_OTEL: ${SIMPLEAUDIT_CHAT_OTLP:-false} + ENABLE_OTEL_TRACES: ${SIMPLEAUDIT_CHAT_OTLP:-false} + OTEL_OTLP_SPAN_EXPORTER: "http" + OTEL_EXPORTER_OTLP_ENDPOINT: ${SIMPLEAUDIT_CHAT_OTLP_ENDPOINT:-http://web:8000} + OTEL_SERVICE_NAME: ${SIMPLEAUDIT_CHAT_OTLP_SERVICE_NAME:-open-webui} + volumes: + - open_webui_data:/app/backend/data + restart: unless-stopped + + chat-proxy: + profiles: ["chat"] + image: ${CADDY_IMAGE:-caddy:2-alpine} + environment: + STUDIO_URL: ${SIMPLEAUDIT_STUDIO_URL:-http://localhost:8000} + SIMPLEAUDIT_CHAT_PROXY_PORT: ${SIMPLEAUDIT_CHAT_PROXY_PORT:-8801} + SIMPLEAUDIT_CHAT_UPSTREAM_PORT: ${SIMPLEAUDIT_CHAT_UPSTREAM_PORT:-8080} + STUDIO_UPSTREAM: web:8000 + volumes: + - ./deploy/compose/Caddyfile.chat:/etc/caddy/Caddyfile:ro + # Served as Open WebUI's /static/custom.css (see Caddyfile.chat) so the + # iframe hides the chat-history sidebar. + - ./chat/embed.css:/etc/caddy/embed/static/custom.css:ro + # Served as Open WebUI's favicon (see Caddyfile.chat). + - ./static/logo.svg:/etc/caddy/branding/logo.svg:ro + ports: + - "${SIMPLEAUDIT_CHAT_PROXY_PORT:-8801}:${SIMPLEAUDIT_CHAT_PROXY_PORT:-8801}" + depends_on: + - open-webui + - web + restart: unless-stopped volumes: postgres_data: + open_webui_data: # Shared Hatchet config; carries the auth-disabled worker JWT from server -> worker. hatchet_config: diff --git a/docs/chat.md b/docs/chat.md new file mode 100644 index 00000000..5c0d6728 --- /dev/null +++ b/docs/chat.md @@ -0,0 +1,272 @@ +# Chat (Open WebUI) + +A module that embeds [Open WebUI](https://openwebui.com) in Studio at `/chat/`, +signed in as the Studio user. It lives in one Django app, `chat/`: + +``` +chat/ + config.py what the module is configured to do, and who you are + views.py urls.py the iframe page and /chat/authz + proxy.py the forward-auth proxy + Open WebUI's lifecycle (embedded) + api.py talking to Open WebUI's API, both directions + management/ sync_chat_models, chat_knowledge + templates/ tests/ +``` + +It is part of the bundle the local one-liner starts +and of the Compose deployment `.env.example` describes, and it can be left out +entirely: with `SIMPLEAUDIT_CHAT` set to `off` (or `disabled`, `false`, `no`, +`0`) or unset in a deployment that does not set it, `/chat/` and `/chat/authz` +return 404, the sidebar has no Chat entry, and nothing extra runs. + +## Why it is an iframe, not a sub-path + +Open WebUI serves from the root of an origin only. It has no base-path setting, +and its HTML references `/static`, `/api` and `/ws` absolutely, so proxying it +under `https://studio/chat/` serves a broken page. It therefore gets its own +origin (a port locally, a host in production) which Studio embeds. + +## Hiding the sidebar in the iframe + +Open WebUI has no embed mode, but its app shell loads `/static/custom.css` on +every page. The proxy answers that one request with [chat/embed.css](../chat/embed.css) +instead of forwarding it, so the iframe renders without the chat-history +sidebar (both its collapsed rail and expanded panel are hidden; the chat fills +the width). The rule lives in +this repo, not in a copy of Open WebUI, so it survives upgrades. Both modes +serve the same file: the Python proxy reads it directly, and the Caddy config +mounts it read-only. To change what the iframe shows, edit `chat/embed.css`. + +## How sign-on works + +Open WebUI's *trusted header* mode: it accepts the identity of whoever calls it +in HTTP headers. A proxy in front of it decides that identity by asking Studio, +using the standard forward-auth contract. + +``` +browser ──► proxy (:8801) + │ strips any client-supplied X-Studio-* header + │ + ├──► Studio GET /chat/authz (browser cookies forwarded) + │ 401 -> proxy redirects the browser to Studio + │ 200 -> X-Studio-Email, X-Studio-Name, X-Studio-Role + │ + └──► Open WebUI (127.0.0.1:8080, no published port) +``` + +Django remains the only authority on identity; nothing outside it reads sessions +or user tables. Open WebUI creates its account on the first request per user and +keeps its own database of chats and settings. + +Role mapping (`X-Studio-Role`, applied on every sign-in): + +| Studio | Open WebUI | +|------------------------------------------|------------| +| superuser, or admin of any workspace | `admin` | +| everyone else | `user` | + +## Security + +**Open WebUI must be reachable only from the proxy.** It believes the headers on +any request it receives, so a client that can connect to it directly can send +`X-Studio-Role: admin` and take over the instance. Embedded mode binds it to +loopback; the compose profile publishes no port for it. Both the Python proxy and +the Caddy config strip client-supplied `X-Studio-*` headers before adding their +own — if you put your own proxy in front, it must do the same. + +## Local (no Docker) + +```bash +uvx simpleaudit-studio # chat is part of the bundle +uvx simpleaudit-studio --disable-chat # leave it out +``` + +Starts Open WebUI (via `open-webui` if installed, otherwise `uvx`) on +127.0.0.1:8080, plus the forward-auth proxy from `infra/chat_proxy.py` on :8801. +Open WebUI's data lives beside Studio's, in `~/.simpleaudit-studio/openwebui/`, +and it runs from that folder so its signing key stays there too. + +The CLI reports what it is doing: it says when chat is starting, warns on a first +run that Open WebUI is being downloaded (~1 GB via `uvx`, a few minutes), prints +where its data and log live, and prints one line when `/chat/` is actually ready +— or why it stopped. Open WebUI's own output goes to `openwebui/server.log`, not +the console. Studio and the worker come up while all this happens. + +`open-webui serve` ignores `HOST`/`PORT` and defaults to **0.0.0.0**:8080, so +Studio passes `--host`/`--port` explicitly. If you override the command with +`SIMPLEAUDIT_CHAT_CMD`, pass those flags yourself — binding it to all interfaces +is what the warning above is about. + +Open WebUI is managed like the embedded Hatchet engine: one instance per Studio +process, started in its own process group, and stopped on the way out — by the +CLI's shutdown and by `atexit`, so Ctrl+C, `kill`, and an unhandled exit all take +it with them. The group matters because `uvx` is only a launcher; signalling it +alone would leave the server running. + +A run that is hard-killed (SIGKILL, a crash, a closed terminal) cannot stop +anything, so its Open WebUI keeps holding the port. The next start finds it +through `openwebui/open-webui.pid` and stops it first — but only when it is a +genuine leftover, i.e. its parent is gone. One that belongs to another running +Studio is left alone, and that start fails on the port instead. + +WebSocket upgrades are tunnelled: the handshake is forwarded with the identity +headers attached, and once Open WebUI answers 101 the two sockets are piped +together — nothing in the proxy understands WebSocket framing. Socket.IO +therefore behaves as it does behind Caddy instead of falling back to polling. + +The proxy keeps no state between requests — an HTTP client with a cookie jar +would hand one browser's session to the next — except a few seconds of "who is +this cookie" (`SIMPLEAUDIT_CHAT_IDENTITY_TTL`, default 5s, 0 to disable). A page +load pulls dozens of assets, and without it each one would ask Django again; a +sign-out takes effect within that window. + +Ollama is switched off (nothing in a Studio deployment serves it). Left on, Open +WebUI polls it on every page load — a failing request in the browser console each +time — and shows an empty Ollama section in its connection settings. The +environment variable only seeds the first start, so a sync also turns it off +through the API. + +## Docker + +Chat is off by default in Compose. Uncomment both switches in `.env` (they are +commented out in `.env.example`) to start chat too: + +```bash +# .env +SIMPLEAUDIT_CHAT=docker # web serves /chat/ +COMPOSE_PROFILES=chat # the two chat containers start +SIMPLEAUDIT_CHAT_URL=http://localhost:8801 # what the browser opens +SIMPLEAUDIT_STUDIO_URL=http://localhost:8000 # where signed-out users are sent + +docker compose up -d +``` + +They are independent, and setting only `SIMPLEAUDIT_CHAT` gives a `/chat/` +page with nothing behind it. + +This runs `open-webui` (no published port) behind `chat-proxy`, a Caddy container +configured by [deploy/compose/Caddyfile.chat](../deploy/compose/Caddyfile.chat). +Studio itself only serves the iframe page and `/chat/authz`. + +On separate hostnames (`studio.example.com` / `chat.example.com`), set +`SESSION_COOKIE_DOMAIN=.example.com` so the proxy receives Studio's session +cookie. Different registrable domains will not work. + +## Settings + +| Variable | Default | Meaning | +|---------------------------------|--------------------------|--------------------------------------------| +| `SIMPLEAUDIT_CHAT` | `embedded` (CLI), unset elsewhere | `embedded`, `docker`, or `off`/`disabled`/`false`/`no`/`0` | +| `SIMPLEAUDIT_CHAT_URL` | `http://localhost:8801` | the origin the iframe loads | +| `SIMPLEAUDIT_CHAT_UPSTREAM` | `http://127.0.0.1:8080` | where Open WebUI listens | +| `SIMPLEAUDIT_CHAT_PROXY_PORT` | `8801` | the proxy's port (both modes) | +| `SIMPLEAUDIT_CHAT_UPSTREAM_PORT`| `8080` | Open WebUI's port (docker mode) | +| `SIMPLEAUDIT_STUDIO_URL` | `http://localhost:8000` | where signed-out users are sent (docker) | +| `SIMPLEAUDIT_CHAT_CMD` | auto | command that starts Open WebUI | +| `SIMPLEAUDIT_CHAT_IDENTITY_TTL` | `5` | seconds the proxy caches who a cookie is | +| `SIMPLEAUDIT_CHAT_SYNC_DELAY` | `2` | seconds a model-connection push waits | +| `SIMPLEAUDIT_CHAT_OTLP` | `false` | export Open WebUI's spans to Studio's OTLP listener | +| `SIMPLEAUDIT_CHAT_OTLP_ENDPOINT`| Studio's web origin | base URL of the OTLP listener (embedded: `http://127.0.0.1:`, docker: `http://web:8000`) | +| `SIMPLEAUDIT_CHAT_OTLP_SERVICE_NAME` | `open-webui` | the service name Open WebUI tags its spans with | + +## Exporting Open WebUI's spans to Studio (OTLP) + +Open WebUI can emit OpenTelemetry traces. When `SIMPLEAUDIT_CHAT_OTLP` is set +to `true`, Studio starts it with tracing enabled and pointed at Studio's own +OTLP listener (`POST /otlp/v1/traces`), so the spans it emits land in the same +place as any other target's — no separate collector needed. + +It is off by default. When enabled it exports **unauthenticated** by default — +no credentials are sent — which matches the listener's default fallback to an +enabled `none` credential. The exporter is the standard OTel one, so the +environment variables are the standard ones, with two Open WebUI specifics: + +- Open WebUI selects the HTTP exporter from `OTEL_OTLP_SPAN_EXPORTER` + (`http`), **not** the standard `OTEL_EXPORTER_OTLP_PROTOCOL`. +- The exporter appends `/v1/traces` to `OTEL_EXPORTER_OTLP_ENDPOINT`, so the + endpoint is the **base** URL, not the full path. + +So enabling it is one line: + +```bash +# .env (docker) — or the equivalent environment in embedded mode +SIMPLEAUDIT_CHAT_OTLP=true +``` + +For an authenticated target (a `basic` or `bearer` OTLP credential instead of a +`none` one), set the matching variables — embedded mode reads them from the +environment it starts Open WebUI with, docker mode passes them straight through: + +```bash +OTEL_BASIC_AUTH_USERNAME=sa_ # from the credential +OTEL_BASIC_AUTH_PASSWORD= # shown once at creation +``` + +If no enabled `none` credential exists and no auth is set, the listener answers +401 and Open WebUI drops the spans. + +## Syncing with Studio (scaffolding) + +`chat/api.py` talks to Open WebUI's API in both directions. It authenticates the +same way the proxy makes the browser authenticate — POST the trusted identity +headers to `/api/v1/auths/signin`, use the token that comes back — so there is no +API key to provision and every call runs as a real Open WebUI user with that +user's role. + +**Push — Studio model connections become Open WebUI providers.** A Studio +connection is a base URL plus a key, which is exactly Open WebUI's +OpenAI-compatible provider config (`OPENAI_API_BASE_URLS` / `OPENAI_API_KEYS` / +`OPENAI_API_CONFIGS`): + +```bash +python manage.py sync_chat_models --dry-run # show what would be pushed +python manage.py sync_chat_models # push every enabled connection +python manage.py sync_chat_models --project demo +``` + +Those lists are also editable by hand in Open WebUI, so each pushed entry carries +a `simpleaudit_connection_id` marker in its config. A sync replaces the marked +entries and leaves everything else where it is — the command says how many of +each. Pushing provider config needs an Open WebUI admin, so the command acts as a +Studio superuser. + +**Pull — what Open WebUI holds.** Knowledge bases come back as plain dicts, so +Studio code never sees Open WebUI's schema: + +```bash +python manage.py chat_knowledge # id, name, file count +python manage.py chat_knowledge --id # one, with its file names +``` + +```python +from chat.api import ChatAPI + +bases = ChatAPI.as_user(request.user).knowledge_bases() +``` + +**The push is automatic.** `chat/signals.py` follows `ModelConnection` and +`RegisteredModel`, so adding a connection, changing a key, disabling one or +registering a model all reach chat on their own. The command stays for a manual +run and for `--dry-run`. + +The push is kept off the request's path. It happens `on_commit`, so Open WebUI +never sees a row that was rolled back; in a background thread, so saving does not +wait on a second service; debounced by `SIMPLEAUDIT_CHAT_SYNC_DELAY` (2s), so an +edit that writes a connection and its models is one push; and best-effort — a +chat that is down or still starting is logged and forgotten, because Studio's own +data is the source of truth. The CLI also syncs once as soon as chat answers, +which covers connections that changed while it was off. + +A connection's registered models become that provider's `model_ids` in Open +WebUI, so chat offers what Studio registered. A connection with no registered +models is left unrestricted. + +## Removing it + +Set `SIMPLEAUDIT_CHAT=disabled`, or pass `--disable-chat` to the CLI. + +To drop the code, delete the `chat/` app and `deploy/compose/Caddyfile.chat`, +then remove its four references: `"chat"` in `INSTALLED_APPS`, the `chat/` route +in `config/urls.py`, the Chat entry in `infra/context_processors.py`, the +`--chat` flag in `simpleaudit_studio/cli.py`, and the `chat` profile in +`docker-compose.yml`. Nothing else refers to it. diff --git a/infra/chat_feature.py b/infra/chat_feature.py new file mode 100644 index 00000000..266b3a09 --- /dev/null +++ b/infra/chat_feature.py @@ -0,0 +1,18 @@ +"""Core-side access to the optional chat module. + +The chat app is optional (SIMPLEAUDIT_CHAT); the core must not import it at +module load, so every core touchpoint goes through this one guarded helper. +""" + + +def chat_enabled() -> bool: + """Whether the chat module is on, without importing it at module load. + + Reads ``chat.config.ENABLED`` at call time (not import time) so runtime + toggles — e.g. tests patching the flag — are honoured. + """ + try: + from chat import config + except ImportError: + return False + return config.ENABLED diff --git a/infra/context_processors.py b/infra/context_processors.py index f5a0817d..a6c078a1 100644 --- a/infra/context_processors.py +++ b/infra/context_processors.py @@ -107,12 +107,14 @@ def nav(request): if user is None or not user.is_authenticated: return {} from accounts.services import is_any_project_admin + from infra.chat_feature import chat_enabled admin = is_any_project_admin(user) path = request.path + entries = _NAV + ((("chat", "Chat", "💬", ("/chat/",), False),) if chat_enabled() else ()) items = [ {"url": reverse(name), "label": label, "icon": icon, "prefixes": prefixes} - for name, label, icon, prefixes, admin_only in _NAV + for name, label, icon, prefixes, admin_only in entries if admin or not admin_only ] diff --git a/infra/engine.py b/infra/engine.py index e866ea02..9b939346 100644 --- a/infra/engine.py +++ b/infra/engine.py @@ -18,6 +18,7 @@ import asyncio import os +from contextlib import nullcontext from typing import Any @@ -314,6 +315,47 @@ def scenario_dict( return scenario +def _collect_evidence_spans( + correlation: Any, provider: Any, *, token_budget: int | None = None +) -> list[dict[str, Any]] | None: + """Collect + select evidence spans for the traces a run recorded. + + The engine records ``turn_id -> trace_id`` in ``correlation`` as it + propagates the W3C ``traceparent``. After the run we fetch the spans for + every recorded trace id from ``provider`` (the provider owns its backend's + eventual consistency — it retries until the trace is available or its + timeout elapses), de-duplicate, and select the evidence-relevant kinds via + the engine's ``select_spans``. + + Returns ``None`` when there is no trace evidence — the caller then judges + on the conversation alone (the normal black-box path). A fetch failure is + logged, never raised: tracing is best-effort evidence and must not fail an + audit run. + """ + if correlation is None or provider is None: + return None + from simpleaudit.tracing.selection import select_spans + + all_spans: list[dict[str, Any]] = [] + seen: set[str] = set() + for tid in correlation.all_trace_ids(): + try: + spans = provider.fetch(tid) + except Exception as exc: # noqa: BLE001 - tracing must not break the run + logger = __import__("logging").getLogger("simpleaudit.engine") + logger.warning("Trace fetch failed for %s: %s", tid, exc) + continue + for span in spans or []: + sid = span.get("span_id") + if sid in seen: + continue + seen.add(sid) + all_spans.append(span) + if not all_spans: + return None + return select_spans(all_spans, token_budget=token_budget).selected or None + + def run_scenario( *, name: str, @@ -330,19 +372,30 @@ def run_scenario( file_uri=None, category: str = "", metadata: dict | None = None, + trace_config: dict | None = None, ) -> dict[str, Any]: """Execute one scenario through the real engine and return a serializable result. - Runs the async ``ModelAuditor.run_scenario`` to completion and returns - ``AuditResult.to_dict()`` plus the language used. Raises ``EngineError`` on - load failure; a mid-conversation/judging failure is captured by the engine - itself as a severity of ``ERROR`` in the returned dict (not raised), matching - the engine's own error-handling contract. + Runs the async ``ModelAuditor.run_async`` (single scenario) to completion + and returns ``AuditResult.to_dict()`` plus the language used. Raises + ``EngineError`` on load failure; a mid-conversation/judging failure is + captured by the engine itself as a severity of ``ERROR`` in the returned + dict (not raised), matching the engine's own error-handling contract. If ``on_turn`` is provided, it is called at each phase boundary with - ``(turn_index, max_turns, role)`` where role is "auditor", "target", or "judge". - NOTE: on_turn is called from within the asyncio event loop — do NOT perform - blocking I/O (e.g. Django ORM) inside it. + ``(turn_index, max_turns, role)`` where role is "auditor", "target", or + "judge". NOTE: on_turn is called from within the asyncio event loop — do + NOT perform blocking I/O (e.g. Django ORM) inside it. + + Tracing (best-effort, Promptfoo parity): when ``trace_config`` is given, a + provider is built from it (``builtin`` OTLP receiver or ``tempo`` fetch) and + a fresh ``TraceCorrelation`` is passed to the engine, which propagates a W3C + ``traceparent`` per turn and records ``turn_id -> trace_id``. After the run + the spans for the recorded trace ids are fetched from the provider, + selected, and attached to the result under ``judgment["evidence_spans"]``. + The engine judges on the conversation (trace evidence is attached post-run + for findings / a trace-aware judge, not re-injected into the judge prompt + here). """ auditor_instance, language = build_model_auditor( target=target, auditor=auditor, judge=judge, generation=generation @@ -352,22 +405,51 @@ def run_scenario( severity_ceiling=severity_ceiling, documents=documents, file_uri=file_uri, category=category, metadata=metadata, ) + + provider = None + correlation = None + audit_run_id = None + if trace_config: + from simpleaudit.tracing.context import TraceCorrelation, new_trace_id + + from infra.tracing import build_trace_provider + + provider = build_trace_provider(trace_config) + if provider is not None: + audit_run_id = f"audit_{new_trace_id()[:12]}" + correlation = TraceCorrelation(audit_run_id=audit_run_id) + try: - # run_async maps the scenario dict onto run_scenario (file_uri, - # documents, judge notes, the scenario facts a judge's post-processor - # reads), the same way AuditExperiment does for repetitions. - results = asyncio.run(auditor_instance.run_async([scenario], language=language, on_turn=on_turn)) + with (provider or nullcontext()): + # run_async maps the scenario dict onto run_scenario (file_uri, + # documents, judge notes, the scenario facts a judge's + # post-processor reads). When tracing, the engine propagates the + # traceparent per turn and records the trace ids in correlation. + results = asyncio.run( + auditor_instance.run_async( + [scenario], + language=language, + on_turn=on_turn, + audit_run_id=audit_run_id, + trace_correlation=correlation, + ) + ) + # Collect evidence spans while the provider is still alive. + evidence_spans = _collect_evidence_spans(correlation, provider) except Exception as exc: raise EngineError(f"Scenario execution crashed: {type(exc).__name__}: {exc}") from exc payload = results[0].to_dict() + if evidence_spans: + judgment = payload.get("judgment") + if not isinstance(judgment, dict): + judgment = {} + judgment["evidence_spans"] = evidence_spans + payload["judgment"] = judgment payload["_language"] = language return payload -_SEV_RANK = {"ERROR": 6, "critical": 5, "high": 4, "medium": 3, "low": 2, "pass": 1} - - def run_scenario_repeated( *, name: str, @@ -388,6 +470,7 @@ def run_scenario_repeated( file_uri=None, category: str = "", metadata: dict | None = None, + trace_config: dict | None = None, ) -> dict[str, Any]: """Execute one scenario N times using AuditExperiment.run_scenario_reps(). @@ -406,6 +489,14 @@ def run_scenario_repeated( If ``on_rep_done`` is provided it is called after each rep with ``(rep_index, rep_result_dict)`` — useful for emitting progress events. If ``cancel_event`` is set, remaining reps are skipped. + + Tracing (best-effort, Promptfoo parity): when ``trace_config`` is given, a + provider is built from it and a fresh ``TraceCorrelation`` is shared across + all reps. Each rep records its ``turn_id -> trace_id`` links, and after the + run the spans for every recorded trace id are fetched, selected, and + attached to each rep's result under ``judgment["evidence_spans"]`` (de- + duplicated per rep). The engine judges on the conversation; trace evidence + is attached post-run, matching the single-rep path. """ _ensure_engine_available() try: @@ -427,6 +518,23 @@ def run_scenario_repeated( model_entry = {k: v for k, v in kwargs.items() if v is not None} model_entry["label"] = f"{kwargs['model']} (platform)" + # Tracing: build a provider + a per-rep TraceCorrelation. Each rep gets its + # own correlation (assigned at the rep boundary) so its turn->trace links + # are attributable to that rep; evidence is fetched per rep after the run. + provider = None + rep_correlations: dict[int, Any] = {} + if trace_config: + from simpleaudit.tracing.context import new_trace_id + + from infra.tracing import build_trace_provider + + provider = build_trace_provider(trace_config) + if provider is not None: + audit_run_id = f"audit_{new_trace_id()[:12]}" + # Stash on the provider so the rep-boundary callback can mint a + # fresh correlation per rep without extra plumbing. + provider._audit_run_id = audit_run_id + # The engine's on_rep_done receives an AuditResults collection with one # result (single scenario). Collect here and emit AFTER asyncio.run() # returns: the outer callback does Django ORM calls, which cannot run @@ -442,8 +550,15 @@ def _on_rep_done(label: str, rep_index: int, total: int, result) -> None: # ``rep_is_done(label, i)`` is consulted right before rep ``i`` starts — # the only hook that fires at every rep boundary. Use it to signal rep - # starts; never skip. on_rep_started must be async-safe (no ORM calls). + # starts and (when tracing) mint a fresh per-rep TraceCorrelation. Must be + # async-safe (no ORM calls). def _rep_is_done(label: str, rep_index: int) -> bool: + if provider is not None: + from simpleaudit.tracing.context import TraceCorrelation + + rep_correlations[rep_index] = TraceCorrelation( + audit_run_id=getattr(provider, "_audit_run_id", None) + ) if on_rep_started: on_rep_started(rep_index) return False @@ -463,12 +578,35 @@ def _rep_is_done(label: str, rep_index: int) -> bool: except Exception as exc: raise EngineError(f"Failed to construct AuditExperiment: {type(exc).__name__}: {exc}") from exc + # Mutable holder the engine reads each rep: the per-rep correlation + # assigned at the rep boundary. A single shared object can't attribute + # spans to individual reps, so we swap it in per rep. + current_correlation: dict[str, Any] = {"corr": None} + + def _tracing_correlation() -> Any: + return current_correlation["corr"] + try: - results = asyncio.run( - experiment.run_scenario_reps( - model_index=0, scenario=scenario, max_turns=max_turns, language=language, on_turn=on_turn, + with (provider or nullcontext()): + results = asyncio.run( + experiment.run_scenario_reps( + model_index=0, scenario=scenario, max_turns=max_turns, language=language, + on_turn=on_turn, + audit_run_id=getattr(provider, "_audit_run_id", None) if provider else None, + trace_correlation=_tracing_correlation, + ) ) - ) + # Collect per-rep evidence spans while the provider is still alive. + if provider is not None: + for rep in reps: + corr = rep_correlations.get(rep.get("_rep_index")) + evidence = _collect_evidence_spans(corr, provider) + if evidence: + judgment = rep.get("judgment") + if not isinstance(judgment, dict): + judgment = {} + judgment["evidence_spans"] = evidence + rep["judgment"] = judgment except Exception as exc: raise EngineError(f"Scenario execution crashed: {type(exc).__name__}: {exc}") from exc @@ -483,18 +621,18 @@ def _rep_is_done(label: str, rep_index: int) -> bool: for rep in reps: on_rep_done(rep.get("_rep_index", 0), rep) - sev_counts: dict[str, int] = {} - for rep in reps: - sev = rep.get("severity", "") - sev_counts[sev] = sev_counts.get(sev, 0) + 1 - modal_severity = max(sev_counts, key=lambda s: (sev_counts[s], _SEV_RANK.get(s, 0))) if sev_counts else "ERROR" - agreement_rate = sev_counts[modal_severity] / len(reps) if reps else 0.0 + # Delegate the modal/agreement aggregation to the engine so the studio + # and the library share one definition (worst-severity tie-break, ERROR + # handling) instead of each hand-rolling it. + from simpleaudit.repeated_results import aggregate_severities + + agg = aggregate_severities([rep.get("severity", "") for rep in reps]) return { "reps": reps, - "aggregated_severity": modal_severity, - "agreement_rate": round(agreement_rate, 4), - "severity_distribution": sev_counts, + "aggregated_severity": agg["most_common_severity"], + "agreement_rate": round(agg["agreement_rate"], 4), + "severity_distribution": agg["severity_distribution"], "n_repetitions": len(reps), "_language": language, } diff --git a/infra/middleware.py b/infra/middleware.py index 863f87a2..ab896966 100644 --- a/infra/middleware.py +++ b/infra/middleware.py @@ -81,12 +81,17 @@ class CsrfCookieMiddleware(MiddlewareMixin): def process_request(self, request): if hasattr(request, "user") and request.user.is_authenticated: from django.conf import settings - from django.middleware.csrf import _add_new_csrf_cookie + from django.middleware.csrf import get_token # Only set the cookie if it's not already present in the request. existing = request.COOKIES.get(settings.CSRF_COOKIE_NAME) if not existing: - _add_new_csrf_cookie(request) + # get_token() is the public API: it populates + # request.META["CSRF_COOKIE"] (and flags the cookie for update), + # which CsrfViewMiddleware.process_response then reads. Calling + # the private _add_new_csrf_cookie() directly skips that META + # assignment and causes a KeyError in process_response. + get_token(request) class ProjectMiddleware(MiddlewareMixin): @@ -110,4 +115,11 @@ def process_request(self, request): membership = request.user.memberships.select_related("project").first() if membership: request.project = membership.project - request.session["active_project_id"] = request.project.id \ No newline at end of file + request.session["active_project_id"] = request.project.id + if not request.project: + # Backstop: a user with no membership (legacy/orphaned account) + # still lands in the shared Default workspace rather than crashing + # every view that dereferences request.project. + from accounts.services import DEFAULT_PROJECT_SLUG + + request.project = Project.objects.filter(slug=DEFAULT_PROJECT_SLUG).first() \ No newline at end of file diff --git a/infra/minimal_config.py b/infra/minimal_config.py index 4a14f880..ac564c6e 100644 --- a/infra/minimal_config.py +++ b/infra/minimal_config.py @@ -17,7 +17,9 @@ import subprocess import threading import time +from contextlib import contextmanager from pathlib import Path +from types import SimpleNamespace from typing import Any logger = logging.getLogger(__name__) @@ -184,6 +186,31 @@ def _kill_stale_sidecars() -> None: pass +@contextmanager +def _isolated_embedded_process(): + """Keep terminal signals away from the engine until the worker has drained. + + Hatchet SDK 1.41.0 offers no process-session option. Adapt only its module + reference, never subprocess.Popen globally. The SDK retains ownership of + stdin and termination, including its parent-exit cleanup. + """ + from hatchet_sdk import embedded + + original = embedded.subprocess + + def spawn(*args, **kwargs): + kwargs["start_new_session"] = True + return original.Popen(*args, **kwargs) + + embedded.subprocess = SimpleNamespace(**{ + **vars(original), "Popen": spawn, + }) + try: + yield + finally: + embedded.subprocess = original + + def start_embedded_hatchet() -> Any: """Start an embedded Hatchet engine and return the client. @@ -218,7 +245,8 @@ def start_embedded_hatchet() -> Any: ) print("\n⏳ Starting embedded Hatchet engine (first run may take ~15s)...") - client = Hatchet.from_embedded(config) + with _isolated_embedded_process(): + client = Hatchet.from_embedded(config) _embedded_client = client print("✅ Hatchet engine ready.\n") return client diff --git a/infra/runs_table.py b/infra/runs_table.py index edefbca8..7894f4ae 100644 --- a/infra/runs_table.py +++ b/infra/runs_table.py @@ -37,8 +37,22 @@ "created_by": "created_by__username", } -# Preference keys the UI may store (value size is capped). +# Preference keys the UI may store (value size is capped). Optional apps +# (e.g. chat) contribute their own keys via EXTRA_PREFERENCE_KEY_PROVIDERS. PREFERENCE_KEYS = {"dashboard_columns"} +EXTRA_PREFERENCE_KEY_PROVIDERS: list = [] + + +def allowed_preference_keys() -> set: + """Core keys plus whatever the optional apps currently allow. + + Providers are called at request time, so an app's keys follow its own + switch (e.g. chat's key is only accepted while chat is on). + """ + keys = set(PREFERENCE_KEYS) + for provider in EXTRA_PREFERENCE_KEY_PROVIDERS: + keys.update(provider()) + return keys MAX_PREFERENCE_BYTES = 20_000 @@ -200,7 +214,7 @@ def post(self, request): except ValueError: return JsonResponse({"error": "Invalid JSON."}, status=400) key = body.get("key") - if key not in PREFERENCE_KEYS: + if key not in allowed_preference_keys(): return JsonResponse({"error": "Unknown preference."}, status=400) value = body.get("value") if len(json.dumps(value)) > MAX_PREFERENCE_BYTES: diff --git a/infra/tests/test_api_e2e.py b/infra/tests/test_api_e2e.py index f20da70c..4afccf26 100644 --- a/infra/tests/test_api_e2e.py +++ b/infra/tests/test_api_e2e.py @@ -6,6 +6,7 @@ This is the "another developer can run this" integration test that validates the entire user journey without requiring a live worker or Docker stack. """ +from django.test import tag from rest_framework.test import APITestCase from accounts.models import Project, ProjectMembership, User @@ -16,6 +17,7 @@ ) +@tag("slow") class APIE2ETest(APITestCase): """Full user journey via the REST API (no live worker).""" diff --git a/infra/tests/test_connection_sharing.py b/infra/tests/test_connection_sharing.py index ab15bb01..988238d5 100644 --- a/infra/tests/test_connection_sharing.py +++ b/infra/tests/test_connection_sharing.py @@ -144,6 +144,21 @@ def test_consumer_sees_shared_connection_readonly(self): self.assertNotContains(page, f'data-conn-open="{self.conn.id}"') self.assertNotContains(page, f'data-discover="{self.conn.id}"') + def test_chat_icon_links_to_the_chat_handoff(self): + # With chat enabled, each model row offers a chat icon that goes + # through /chat/with// (which validates the model + # server-side). The connection id disambiguates the same model id + # registered under more than one connection. + with mock.patch("chat.config.ENABLED", True): + page = self._page_as(self.owner_user, self.owner_ws) + self.assertContains(page, f'/chat/with/{self.conn.id}/{self.model.model_id}') + self.assertContains(page, "Open a chat with") + + def test_chat_icon_hidden_when_chat_is_disabled(self): + with mock.patch("chat.config.ENABLED", False): + page = self._page_as(self.owner_user, self.owner_ws) + self.assertNotContains(page, "/chat/with/") + def test_owner_sees_edit_controls(self): page = self._page_as(self.owner_user, self.owner_ws) self.assertContains(page, f'data-conn-open="{self.conn.id}"') diff --git a/infra/tests/test_embedded_cleanup.py b/infra/tests/test_embedded_cleanup.py index 12e5951b..df75818c 100644 --- a/infra/tests/test_embedded_cleanup.py +++ b/infra/tests/test_embedded_cleanup.py @@ -68,3 +68,44 @@ def test_postgres_of_a_running_instance_is_left_alone(self): mc._cleanup_stale_embedded_pg(str(self.dir)) kill.assert_not_called() self.assertTrue(self.lock.exists()) + + +class EmbeddedSignalTests(SimpleTestCase): + def test_sidecar_has_separate_process_group_and_explicit_stop(self): + """Exercise SDK spawning with a real child, without starting Postgres.""" + import subprocess + import sys + + from hatchet_sdk import EmbeddedHatchetConfig, embedded + + child = None + real_popen = subprocess.Popen + + def spawn_stub(*args, **kwargs): + nonlocal child + child = real_popen( + [sys.executable, "-c", + "import sys; sys.stdin.buffer.read()"], + **kwargs, + ) + return child + + handshake = embedded.Handshake( + token="test", tenant_id="test", grpc_address="localhost:1", + api_url="http://localhost:1", + ) + try: + with mock.patch.object(subprocess, "Popen", side_effect=spawn_stub), \ + mock.patch.object(embedded, "_wait_for_handshake", return_value=handshake), \ + mc._isolated_embedded_process(): + sidecar = embedded.start_embedded_sidecar( + EmbeddedHatchetConfig(binary_path=sys.executable) + ) + self.assertNotEqual(os.getpgid(child.pid), os.getpgrp()) + self.assertIsNone(child.poll()) + sidecar.stop() + self.assertIsNotNone(child.poll()) + finally: + if child is not None and child.poll() is None: + child.terminate() + child.wait(timeout=5) diff --git a/infra/tests/test_engine_integration.py b/infra/tests/test_engine_integration.py index 6792aa18..f9477d8a 100644 --- a/infra/tests/test_engine_integration.py +++ b/infra/tests/test_engine_integration.py @@ -9,7 +9,7 @@ """ from unittest import mock -from django.test import TestCase +from django.test import TestCase, tag from django.utils import timezone from accounts.models import Project, ProjectMembership, User @@ -70,6 +70,7 @@ def _build_run(user, project): return run, item +@tag("slow") class SecretValidationTest(TestCase): """Fail-fast secret resolution: a missing env var must raise a clean EngineError naming the role + reference, not an opaque any_llm MissingApiKeyError.""" @@ -200,6 +201,79 @@ async def run_scenario_reps(self, **kw): self.assertEqual(entry["judge_kwargs"], {"timeout": 9, "max_retries": 1}) +class RepeatedTracingTest(TestCase): + """run_scenario_repeated forwards tracing and attaches per-rep evidence.""" + + def test_forwards_trace_config_to_engine(self): + from infra import engine + + captured = {} + + class FakeExperiment: + def __init__(self, models, **kw): + pass + + async def run_scenario_reps(self, **kw): + captured["audit_run_id"] = kw.get("audit_run_id") + captured["corr_fn"] = kw.get("trace_correlation") + return [ + type("R", (), {"to_dict": lambda self: {"severity": "pass", "judgment": {}}})() + for _ in range(2) + ] + + # A fake provider (tempo mode) that yields no spans; fetch returns []. + fake_provider = mock.MagicMock() + fake_provider.fetch.return_value = [] + fake_provider._audit_run_id = "audit_x" + + with mock.patch("simpleaudit.experiment.AuditExperiment", FakeExperiment), \ + mock.patch("infra.tracing.build_trace_provider", return_value=fake_provider): + result = engine.run_scenario_repeated( + name="s", description="d", expected_behavior=None, test_prompt=None, + target=_snap("tgt"), auditor=_snap("aud"), judge=_snap("jdg"), + n_repetitions=2, + trace_config={"mode": "tempo", "base_url": "http://tempo.local"}, + ) + + # Tracing params were forwarded to the engine (a fresh audit_run_id is + # minted for the run). + self.assertTrue(str(captured["audit_run_id"]).startswith("audit_")) + self.assertTrue(callable(captured["corr_fn"])) + # Two reps were executed. + self.assertEqual(result["n_repetitions"], 2) + self.assertEqual(len(result["reps"]), 2) + + def test_no_trace_config_is_noop(self): + from infra import engine + + captured = {} + + class FakeExperiment: + def __init__(self, models, **kw): + pass + + async def run_scenario_reps(self, **kw): + captured["audit_run_id"] = kw.get("audit_run_id") + captured["corr_fn"] = kw.get("trace_correlation") + return [ + type("R", (), {"to_dict": lambda self: {"severity": "pass", "judgment": {}}})() + for _ in range(2) + ] + + with mock.patch("simpleaudit.experiment.AuditExperiment", FakeExperiment): + result = engine.run_scenario_repeated( + name="s", description="d", expected_behavior=None, test_prompt=None, + target=_snap("tgt"), auditor=_snap("aud"), judge=_snap("jdg"), + n_repetitions=2, + ) + + self.assertIsNone(captured["audit_run_id"]) + # The correlation callable is always forwarded; with no trace_config it + # resolves to None (no tracing). + self.assertIsNone(captured["corr_fn"]()) + self.assertEqual(result["n_repetitions"], 2) + + class ProviderNormalizationTest(TestCase): def test_known_provider_passthrough(self): from infra.engine import _normalize_provider diff --git a/infra/tests/test_experiments.py b/infra/tests/test_experiments.py index 27c6c9d2..aa724b25 100644 --- a/infra/tests/test_experiments.py +++ b/infra/tests/test_experiments.py @@ -7,7 +7,7 @@ from datetime import timedelta from unittest import mock -from django.test import Client, TestCase +from django.test import Client, TestCase, tag from django.utils import timezone from audits.models import AuditRun, Experiment @@ -26,6 +26,7 @@ _PROVENANCE = mock.Mock(version="0.2.1", commit="", source="metadata") +@tag("slow") class _ExperimentBase(TestCase): def setUp(self): self.user = UserFactory() diff --git a/infra/tests/test_judges.py b/infra/tests/test_judges.py index 798a3406..17fb9e9b 100644 --- a/infra/tests/test_judges.py +++ b/infra/tests/test_judges.py @@ -3,7 +3,7 @@ Run: SIMPLEAUDIT_LOCAL_SQLITE=1 uv run manage.py test infra.tests.test_judges """ -from django.test import Client, TestCase +from django.test import Client, TestCase, tag from infra.tests.factories import ( AuditRunFactory, @@ -25,6 +25,7 @@ ) +@tag("slow") class _Base(TestCase): role = "admin" diff --git a/infra/tests/test_minimal_config.py b/infra/tests/test_minimal_config.py index e7fa88c4..880e9e89 100644 --- a/infra/tests/test_minimal_config.py +++ b/infra/tests/test_minimal_config.py @@ -13,7 +13,10 @@ import urllib.request from unittest.mock import patch -from django.test import TestCase, tag +from django.contrib.auth import get_user_model +from django.test import TestCase, override_settings, tag + +User = get_user_model() class TestDemoBootSequence(TestCase): @@ -157,3 +160,70 @@ def test_get_client_returns_embedded(self): self.assertIs(c, self._hatchet_client) finally: w._CLIENT = None + + +class TestAutoLoginToken(TestCase): + """The /auto-login/ endpoint must be one-time: a single-use token URL. + + The demo server binds 0.0.0.0, so the token is what keeps other machines + on the LAN from signing in as the bootstrap user. + """ + + TOKEN = "test-one-time-token-123" + + def setUp(self): + self.user = User.objects.create_user(username="studio", password="x") + + def tearDown(self): + os.environ.pop("SIMPLEAUDIT_AUTO_LOGIN_TOKEN", None) + + def _login(self, token=None): + url = "/auto-login/" + if token is not None: + url += f"?token={token}" + # follow=True: the view 302s to the dashboard after signing in. + return self.client.get(url, follow=True) + + @override_settings(MINIMAL_CONFIG=True) + def test_valid_token_logs_in_and_consumes(self): + with patch.dict(os.environ, {"SIMPLEAUDIT_AUTO_LOGIN_TOKEN": self.TOKEN}): + response = self._login(self.TOKEN) + self.assertEqual(response.status_code, 200) + self.assertEqual(response.wsgi_request.user.username, "studio") + # The token is single-use: it must be gone after the first login. + self.assertNotIn("SIMPLEAUDIT_AUTO_LOGIN_TOKEN", os.environ) + + @override_settings(MINIMAL_CONFIG=True) + def test_token_cannot_be_replayed(self): + with patch.dict(os.environ, {"SIMPLEAUDIT_AUTO_LOGIN_TOKEN": self.TOKEN}): + first = self._login(self.TOKEN) + self.assertEqual(first.status_code, 200) + # Second use of the same URL: token already consumed -> 404. + second = self._login(self.TOKEN) + self.assertEqual(second.status_code, 404) + + @override_settings(MINIMAL_CONFIG=True) + def test_invalid_token_rejected(self): + with patch.dict(os.environ, {"SIMPLEAUDIT_AUTO_LOGIN_TOKEN": self.TOKEN}): + response = self._login("wrong-token") + self.assertEqual(response.status_code, 404) + # Token must survive a failed attempt so the real URL still works. + self.assertEqual(os.environ.get("SIMPLEAUDIT_AUTO_LOGIN_TOKEN"), self.TOKEN) + + @override_settings(MINIMAL_CONFIG=True) + def test_missing_token_rejected(self): + with patch.dict(os.environ, {"SIMPLEAUDIT_AUTO_LOGIN_TOKEN": self.TOKEN}): + response = self._login() + self.assertEqual(response.status_code, 404) + + @override_settings(MINIMAL_CONFIG=True) + def test_no_token_configured_rejects_everything(self): + os.environ.pop("SIMPLEAUDIT_AUTO_LOGIN_TOKEN", None) + response = self._login(self.TOKEN) + self.assertEqual(response.status_code, 404) + + @override_settings(MINIMAL_CONFIG=False) + def test_disabled_outside_minimal_config(self): + with patch.dict(os.environ, {"SIMPLEAUDIT_AUTO_LOGIN_TOKEN": self.TOKEN}): + response = self._login(self.TOKEN) + self.assertEqual(response.status_code, 404) diff --git a/infra/tests/test_monitors.py b/infra/tests/test_monitors.py index 801fb31a..8d6b164d 100644 --- a/infra/tests/test_monitors.py +++ b/infra/tests/test_monitors.py @@ -4,15 +4,17 @@ SIMPLEAUDIT_LOCAL_SQLITE=1 uv run manage.py test infra.tests.test_monitors """ from datetime import UTC, timedelta -from unittest import mock +from unittest import mock, skipUnless -from django.test import Client, TestCase +from django.db import connection +from django.test import Client, TestCase, tag from django.utils import timezone from audits.models import AuditRun, Monitor from audits.monitors import ( advance, drift_series, + due_monitors, run_due_monitors, two_proportion_z, wilson, @@ -40,6 +42,7 @@ def _result(severities): return {"reps": [{"severity": s} for s in severities], "n_repetitions": len(severities)} +@tag("slow") class MonitorTestBase(TestCase): def setUp(self): self.user = UserFactory() @@ -524,3 +527,36 @@ def test_experiment_repeat_creates_linked_monitors(self): self.assertEqual(len(page.context["current_runs"]), 2) pooled = self.client.get(f"/experiments/{exp.id}/?pool=1") self.assertTrue(pooled.context["pool"]) + + +class DueMonitorQueryTests(MonitorTestBase): + """The claim query must lock only the monitor rows. + + ``select_related`` pulls in ``last_run`` and ``created_by``, both nullable, so + the join is a LEFT OUTER JOIN. PostgreSQL rejects ``FOR UPDATE`` against the + nullable side of one ("FOR UPDATE cannot be applied to the nullable side of an + outer join"), which made every tick fail on a production database while the + sweeper logged the error and carried on. + + SQLite drops ``FOR UPDATE`` altogether, so the failure cannot be reproduced on + the suite's default backend. These assertions are on the query the ORM builds, + which is backend-independent; the execution test below runs on PostgreSQL only. + """ + + def test_claim_query_locks_only_monitor_rows(self): + query = due_monitors(timezone.now()).query + self.assertEqual(query.select_for_update_of, ("self",)) + self.assertTrue(query.select_for_update_skip_locked) + + def test_claim_query_still_joins_the_nullable_relations(self): + # If these stop being selected the `of` above is no longer load-bearing, + # and a future change could drop it without any test noticing. + self.assertEqual( + set(due_monitors(timezone.now()).query.select_related), + {"last_run", "project", "created_by"}, + ) + + @skipUnless(connection.vendor == "postgresql", "outer-join locking is PostgreSQL-specific") + def test_tick_executes_on_postgresql(self): + # The regression itself: this raises NotSupportedError without `of`. + run_due_monitors() diff --git a/infra/tests/test_tracing.py b/infra/tests/test_tracing.py new file mode 100644 index 00000000..4829c29f --- /dev/null +++ b/infra/tests/test_tracing.py @@ -0,0 +1,248 @@ +"""Trace acquisition: the Tempo provider, the factory, and the engine hand-off. + +The tracing core lives in the SimpleAudit engine (``simpleaudit.tracing``). +These tests cover the studio-side layer (``infra.tracing``) and how +``infra.engine.run_scenario`` wires the engine's ``TraceCorrelation`` + +provider into a run. + +Run: + SIMPLEAUDIT_LOCAL_SQLITE=1 uv run manage.py test infra.tests.test_tracing +""" +from unittest import mock + +from django.test import TestCase + +from infra.tracing import TempoTraceProvider, build_trace_provider + + +def _otlp_payload(trace_id: str, spans: list[dict]) -> dict: + """Wrap normalized-ish spans in an OTLP/HTTP JSON ExportTraceServiceRequest.""" + resource_spans = [] + for s in spans: + attrs = [{"key": k, "value": {"stringValue": str(v)}} for k, v in (s.get("attributes") or {}).items()] + resource_spans.append( + { + "resource": {"attributes": []}, + "scopeSpans": [ + { + "spans": [ + { + "traceId": trace_id, + "spanId": s["span_id"], + "parentSpanId": s.get("parent_span_id") or "", + "name": s.get("name", "span"), + "kind": 1, + "startTimeUnixNano": 0, + "endTimeUnixNano": 1, + "attributes": attrs, + "status": {"code": 0}, + } + ] + } + ], + } + ) + return {"resourceSpans": resource_spans} + + +def _tempo_handler(payload_by_trace: dict): + """Build an httpx.MockTransport handler that serves Tempo's trace API. + + ``payload_by_trace`` maps trace_id -> OTLP JSON body. A trace not present + (or mapped to ``None``) yields a 404, simulating Tempo's eventual + consistency (not yet queryable). + """ + import httpx + + def handler(request: httpx.Request) -> httpx.Response: + trace_id = request.url.path.rsplit("/", 1)[-1] + if trace_id not in payload_by_trace or payload_by_trace[trace_id] is None: + return httpx.Response(404, json={"error": "trace not found"}) + return httpx.Response(200, json=payload_by_trace[trace_id]) + + return httpx.MockTransport(handler) + + +class TempoProviderTest(TestCase): + """TempoTraceProvider fetches + parses a trace from Tempo's API.""" + + def _provider(self, transport, **kw): + return TempoTraceProvider("http://tempo.local", transport=transport, retry_interval=0.01, **kw) + + def test_fetch_returns_normalized_spans(self): + tid = "a" * 32 + payload = _otlp_payload(tid, [ + {"span_id": "1" * 16, "name": "retrieve", "attributes": {"openinference.span.kind": "RETRIEVER"}}, + {"span_id": "2" * 16, "parent_span_id": "1" * 16, "name": "llm", "attributes": {}}, + ]) + provider = self._provider(_tempo_handler({tid: payload})) + spans = provider.fetch(tid) + self.assertEqual(len(spans), 2) + by_name = {s["name"]: s for s in spans} + self.assertEqual(by_name["retrieve"]["kind"], "RETRIEVER") + self.assertEqual(by_name["retrieve"]["trace_id"], tid) + self.assertEqual(by_name["llm"]["parent_span_id"], "1" * 16) + + def test_fetch_retries_until_available(self): + """A trace that 404s then appears is returned after a retry.""" + tid = "b" * 32 + payload = _otlp_payload(tid, [{"span_id": "3" * 16, "name": "tool", "attributes": {}}]) + state = {"n": 0} + + import httpx + + def handler(request: httpx.Request) -> httpx.Response: + state["n"] += 1 + if state["n"] == 1: + return httpx.Response(404, json={"error": "not yet"}) + return httpx.Response(200, json=payload) + + provider = self._provider(httpx.MockTransport(handler), timeout=2.0) + spans = provider.fetch(tid) + self.assertEqual(len(spans), 1) + self.assertGreaterEqual(state["n"], 2) + + def test_fetch_returns_empty_on_timeout(self): + tid = "c" * 32 + provider = self._provider(_tempo_handler({}), timeout=0.05) + self.assertEqual(provider.fetch(tid), []) + + def test_fetch_returns_empty_on_http_error(self): + tid = "d" * 32 + import httpx + + def handler(request: httpx.Request) -> httpx.Response: + return httpx.Response(500, json={"error": "boom"}) + + provider = self._provider(httpx.MockTransport(handler), timeout=0.05) + self.assertEqual(provider.fetch(tid), []) + + def test_empty_base_url_returns_empty(self): + provider = TempoTraceProvider("", transport=_tempo_handler({})) + self.assertEqual(provider.fetch("e" * 32), []) + + def test_context_manager_is_noop(self): + provider = self._provider(_tempo_handler({})) + with provider as p: + self.assertIs(p, provider) + self.assertIsNone(provider.endpoint) + + +class BuildTraceProviderTest(TestCase): + """build_trace_provider maps a config dict onto the right provider.""" + + def test_none_or_empty_config_returns_none(self): + self.assertIsNone(build_trace_provider(None)) + self.assertIsNone(build_trace_provider({})) + + def test_none_mode_returns_none(self): + self.assertIsNone(build_trace_provider({"mode": "none"})) + self.assertIsNone(build_trace_provider({"mode": "off"})) + + def test_builtin_mode_returns_builtin_otlp(self): + from simpleaudit.tracing.provider import BuiltinOTLP + + provider = build_trace_provider({"mode": "builtin"}) + self.assertIsInstance(provider, BuiltinOTLP) + + def test_default_mode_is_builtin(self): + from simpleaudit.tracing.provider import BuiltinOTLP + + self.assertIsInstance(build_trace_provider({"base_url": "http://x"}), BuiltinOTLP) + + def test_tempo_mode_returns_tempo_provider(self): + provider = build_trace_provider({"mode": "tempo", "base_url": "http://tempo.local"}) + self.assertIsInstance(provider, TempoTraceProvider) + + def test_tempo_mode_requires_base_url(self): + with self.assertRaises(ValueError): + build_trace_provider({"mode": "tempo"}) + + def test_unknown_mode_raises(self): + with self.assertRaises(ValueError): + build_trace_provider({"mode": "jaeger"}) + + +class RunScenarioTracingWiringTest(TestCase): + """run_scenario passes the correlation to the engine and attaches evidence.""" + + def _snap(self, model_id, **extra): + return {"model_id": model_id, "provider": "openai", "base_url": f"http://{model_id}.local/v1", **extra} + + def test_trace_config_records_correlation_and_attaches_evidence(self): + from types import SimpleNamespace + + from infra.engine import run_scenario + + tid = "f" * 32 + payload = _otlp_payload(tid, [ + {"span_id": "4" * 16, "name": "retrieve", "attributes": {"openinference.span.kind": "RETRIEVER"}}, + ]) + transport = _tempo_handler({tid: payload}) + + captured = {} + + class FakeAuditor: + async def run_async(self, scenarios, **kwargs): + captured["kwargs"] = kwargs + correlation = kwargs.get("trace_correlation") + if correlation is not None: + # Simulate the engine recording the scenario's trace id. + correlation.record("scen_t1", tid) + result = SimpleNamespace(to_dict=lambda: {"severity": "pass", "judgment": {"severity": "pass"}}) + return [result] + + with mock.patch("infra.engine.build_model_auditor", return_value=(FakeAuditor(), "English")): + payload = run_scenario( + name="s", description="d", expected_behavior=None, test_prompt=None, + target=self._snap("t"), auditor=self._snap("a"), judge=self._snap("j"), + generation={"max_turns": 1}, + trace_config={ + "mode": "tempo", + "base_url": "http://tempo.local", + "transport": transport, + "retry_interval": 0.01, + }, + ) + + # The engine got a correlation + audit_run_id. + self.assertIsNotNone(captured["kwargs"].get("trace_correlation")) + self.assertTrue(captured["kwargs"].get("audit_run_id")) + # Evidence was fetched from the provider and attached to the judgment. + self.assertIn("evidence_spans", payload["judgment"]) + self.assertEqual(payload["judgment"]["evidence_spans"][0]["name"], "retrieve") + + def test_no_trace_config_leaves_judgment_untouched(self): + from types import SimpleNamespace + + from infra.engine import run_scenario + + captured = {} + + class FakeAuditor: + async def run_async(self, scenarios, **kwargs): + captured["kwargs"] = kwargs + result = SimpleNamespace(to_dict=lambda: {"severity": "pass", "judgment": {"severity": "pass"}}) + return [result] + + with mock.patch("infra.engine.build_model_auditor", return_value=(FakeAuditor(), "English")): + payload = run_scenario( + name="s", description="d", expected_behavior=None, test_prompt=None, + target=self._snap("t"), auditor=self._snap("a"), judge=self._snap("j"), + generation={"max_turns": 1}, + ) + + self.assertIsNone(captured["kwargs"].get("trace_correlation")) + self.assertIsNone(captured["kwargs"].get("audit_run_id")) + self.assertNotIn("evidence_spans", payload["judgment"]) + + def test_collect_evidence_spans_none_when_no_traces(self): + from infra.engine import _collect_evidence_spans + + self.assertIsNone(_collect_evidence_spans(None, None)) + + class _NoTraces: + def all_trace_ids(self): + return [] + + self.assertIsNone(_collect_evidence_spans(_NoTraces(), object())) diff --git a/infra/tests/test_user_project_assignment.py b/infra/tests/test_user_project_assignment.py new file mode 100644 index 00000000..b9fa30e2 --- /dev/null +++ b/infra/tests/test_user_project_assignment.py @@ -0,0 +1,112 @@ +"""Every user-creation path must land the new user in the 'default' workspace. + +A user with no project membership makes ``request.project`` resolve to ``None`` +and 500s every project-scoped view. These tests pin the invariant that each +door into the system (admin add-user, self-registration, demo signup) grants a +viewer membership in the shared Default workspace, and that the middleware +backstop keeps even a membership-less user from crashing. +""" +from django.test import Client, TestCase + +from accounts.models import Project, ProjectMembership, User +from infra.tests.factories import ProjectFactory +from infra.tests.utils import login, post_json, superuser + + +def _default_project(): + return Project.objects.create(name="Default", slug="default") + + +class AdminCreateUserAssignmentTest(TestCase): + def setUp(self): + self.client = Client(SERVER_NAME="localhost") + self.admin = superuser() + login(self.client, self.admin) + + def test_admin_added_user_gets_default_viewer(self): + _default_project() + resp = post_json( + self.client, + "/api/admin/users/", + {"username": "newbie", "email": "newbie@test.com", "password": "Str0ng-pass-123"}, + ) + self.assertEqual(resp.status_code, 201) + user = User.objects.get(username="newbie") + default = Project.objects.get(slug="default") + self.assertTrue( + ProjectMembership.objects.filter( + project=default, user=user, role=ProjectMembership.Role.VIEWER + ).exists() + ) + + def test_admin_added_user_without_default_project_is_safe(self): + # No 'default' project exists (e.g. pre-bootstrap): user is created, + # just without a membership — no crash. + resp = post_json( + self.client, + "/api/admin/users/", + {"username": "orphan", "password": "Str0ng-pass-123"}, + ) + self.assertEqual(resp.status_code, 201) + self.assertEqual(ProjectMembership.objects.filter(user__username="orphan").count(), 0) + + +class RegisterApiAssignmentTest(TestCase): + def test_register_grants_default_viewer(self): + _default_project() + resp = self.client.post( + "/api/auth/register/", + {"username": "carol", "email": "carol@example.com", "password": "another-strong-pass"}, + ) + self.assertEqual(resp.status_code, 201) + user = User.objects.get(username="carol") + default = Project.objects.get(slug="default") + self.assertTrue( + ProjectMembership.objects.filter( + project=default, user=user, role=ProjectMembership.Role.VIEWER + ).exists() + ) + + +class DemoSignupAssignmentTest(TestCase): + def test_demo_register_grants_default_viewer(self): + _default_project() + resp = self.client.post( + "/register/", + {"username": "demo-user", "password": "Str0ng-pass-123", "email": "demo@test.com"}, + ) + self.assertEqual(resp.status_code, 302) + user = User.objects.get(username="demo-user") + default = Project.objects.get(slug="default") + self.assertTrue( + ProjectMembership.objects.filter( + project=default, user=user, role=ProjectMembership.Role.VIEWER + ).exists() + ) + + +class ProjectMiddlewareBackstopTest(TestCase): + """A user with no membership must still resolve a project (no 500).""" + + def test_membershipless_user_falls_back_to_default(self): + default = _default_project() + # A non-default project exists too, to prove the fallback is by slug. + ProjectFactory(name="Other", slug="other") + user = User.objects.create_user(username="orphan", password="orphan-pass-1") + + client = Client(SERVER_NAME="localhost") + client.force_login(user) + resp = client.get("/") + self.assertEqual(resp.status_code, 200) + self.assertEqual(resp.wsgi_request.project, default) + + def test_membershipless_user_without_default_project_gets_none(self): + # No 'default' project at all: request.project stays None (the views + # must tolerate that), but the request itself must not 500. + ProjectFactory(name="Other", slug="other") + user = User.objects.create_user(username="orphan2", password="orphan-pass-2") + + client = Client(SERVER_NAME="localhost") + client.force_login(user) + resp = client.get("/healthz") + self.assertEqual(resp.status_code, 200) diff --git a/infra/tracing.py b/infra/tracing.py new file mode 100644 index 00000000..9458d0e2 --- /dev/null +++ b/infra/tracing.py @@ -0,0 +1,176 @@ +"""Trace acquisition for audits — the studio-side layer over the engine's tracing API. + +The SimpleAudit engine (``simpleaudit.tracing``) owns the core: the W3C +``traceparent`` correlation, the OTLP receiver, the span store, and the two +provider base classes. This module is the thin studio layer that: + +* builds a provider from a run's trace config (``build_trace_provider``), and +* implements the **external** provider for Grafana Tempo (``TempoTraceProvider``), + which the engine ships only as the abstract ``ExternalTraceProvider`` base. + +Two acquisition modes (Promptfoo parity): + +1. ``builtin`` — the engine's ``BuiltinOTLP`` ephemeral OTLP receiver. For + controlled/staging targets you can point at ``provider.endpoint``. +2. ``tempo`` — ``TempoTraceProvider``. For production targets that already ship + traces to the owner's Tempo backend; we fetch the matching trace by id. + +The engine propagates the ``traceparent`` during the run and records +``turn_id -> trace_id`` in a ``TraceCorrelation``; after the run the caller +fetches spans for the recorded trace ids and hands the selected evidence to the +judge. See ``infra.engine.run_scenario``. +""" +from __future__ import annotations + +import logging +import time +from typing import Any, Self + +logger = logging.getLogger("simpleaudit.tracing") + + +class TempoTraceProvider: + """Fetch traces from a Grafana Tempo backend by trace id. + + Subclasses the engine's ``ExternalTraceProvider`` contract (``start`` / + ``stop`` / ``fetch``) and implements ``_fetch_remote`` against Tempo's + trace API (``GET {base_url}/api/traces/{trace_id}``), which returns an + OTLP/HTTP JSON ``ExportTraceServiceRequest``. The response is parsed with + the engine's ``parse_otlp_json`` and normalized via ``SpanStore`` so the + spans match the schema the judge's evidence selection expects. + + Tempo is eventually consistent: a trace is not queryable until its spans + have been flushed. ``fetch`` therefore retries until the trace appears or + ``timeout`` elapses, sleeping ``retry_interval`` between attempts. + """ + + def __init__( + self, + base_url: str, + *, + timeout: float = 15.0, + retry_interval: float = 0.5, + headers: dict[str, str] | None = None, + transport: Any = None, + ) -> None: + self._base_url = (base_url or "").rstrip("/") + self._timeout = timeout + self._retry_interval = retry_interval + self._headers = headers or {} + # Injectable for tests (httpx.MockTransport); None = real network. + self._transport = transport + + # -- TraceProvider contract (matches simpleaudit.tracing.TraceProvider) -- + + def start(self) -> TempoTraceProvider: + return self + + def stop(self) -> None: + return None + + @property + def endpoint(self) -> str | None: + # Fetch-based: the target already exports to its own backend. + return None + + def fetch(self, trace_id: str) -> list[dict[str, Any]]: + return self._fetch_remote(trace_id) + + def __enter__(self) -> Self: + return self.start() + + def __exit__(self, *exc: object) -> None: + self.stop() + + # -- Tempo specifics ----------------------------------------------------- + + def _fetch_remote(self, trace_id: str) -> list[dict[str, Any]]: + """GET the trace, retrying until available or the timeout elapses.""" + if not self._base_url: + return [] + url = f"{self._base_url}/api/traces/{trace_id}" + deadline = time.monotonic() + self._timeout + while True: + spans = self._fetch_once(url) + if spans: + return spans + if time.monotonic() >= deadline: + return [] + time.sleep(self._retry_interval) + + def _fetch_once(self, url: str) -> list[dict[str, Any]]: + """One GET + parse. Returns [] when the trace is not (yet) available. + + A 404 is Tempo's "not yet queryable" (eventual consistency) — the + caller retries. Any other failure is logged and treated as no spans: + tracing is best-effort evidence and must not break the run. + """ + import httpx + from simpleaudit.tracing.otlp import parse_otlp_json + from simpleaudit.tracing.store import SpanStore + + try: + client = httpx.Client( + headers=self._headers, + timeout=self._timeout, + transport=self._transport, + ) + with client: + resp = client.get(url) + if resp.status_code == 404: + return [] + resp.raise_for_status() + payload = resp.json() + except Exception as exc: # noqa: BLE001 - tracing must not break the run + logger.warning("Tempo trace fetch failed for %s: %s", url, exc) + return [] + + raw_spans = parse_otlp_json(payload) + if not raw_spans: + return [] + store = SpanStore() + store.add_many(raw_spans) + return store.all() + + +def build_trace_provider(config: dict[str, Any] | None) -> Any | None: + """Build a trace provider from a run's trace config, or ``None``. + + ``config`` is a dict (typically from an AuditRun's trace settings) with: + + * ``mode`` — ``"builtin"`` or ``"tempo"`` (default ``"builtin"``). + * ``base_url`` — Tempo base URL (required for ``tempo`` mode). + * ``timeout`` / ``retry_interval`` / ``headers`` — Tempo fetch tuning. + + Returns ``None`` when tracing is disabled (``config`` is falsy or + ``mode`` is ``"none"``/``"off"``). The returned provider is NOT started; + the caller manages its lifecycle (``with provider:`` or explicit + ``start()``/``stop()``). + """ + if not config: + return None + mode = str(config.get("mode") or "builtin").strip().lower() + if mode in ("none", "off", "disabled", ""): + return None + + if mode == "tempo": + base_url = (config.get("base_url") or "").strip() + if not base_url: + raise ValueError("trace config mode='tempo' requires a non-empty base_url") + return TempoTraceProvider( + base_url, + timeout=float(config.get("timeout") or 15.0), + retry_interval=float(config.get("retry_interval") or 0.5), + headers=config.get("headers") or None, + transport=config.get("transport"), + ) + + if mode == "builtin": + from simpleaudit.tracing.provider import BuiltinOTLP + + return BuiltinOTLP( + host=config.get("host") or "127.0.0.1", + port=int(config.get("port") or 0), + ) + + raise ValueError(f"Unknown trace mode: {mode!r} (expected 'builtin' or 'tempo')") diff --git a/infra/ui.py b/infra/ui.py index c597b398..a4c45917 100644 --- a/infra/ui.py +++ b/infra/ui.py @@ -1,6 +1,7 @@ """Server-rendered UI — Django CBVs + Forms + HTMX.""" import csv import hashlib +import hmac import io import itertools import json @@ -24,6 +25,7 @@ from audits.events import ScenarioResult from audits.models import AuditRun from audits.services import create_audit_run, frozen_name, submit_audit_run +from infra.chat_feature import chat_enabled from infra.hashing import scenario_revision_hash from scenarios.models import ( Scenario, @@ -171,6 +173,9 @@ def post(self, request, *args, **kwargs): error = "Username already taken." else: user = User.objects.create_user(username=username, password=password, email=email) + from accounts.services import grant_default_project + + grant_default_project(user) login(request, user) return redirect("dashboard") return self.render_to_response(self.get_context_data(error=error)) @@ -182,10 +187,12 @@ def logout_view(request): def auto_login_view(request): - """One-click sign-in for the local one-liner demo (`uvx simpleaudit-studio`). + """One-click, one-time sign-in for the local one-liner demo. - The CLI opens this URL in the default browser after startup; it logs the - visitor in as the shared bootstrap user and lands them on the dashboard. + The CLI generates a single-use token at startup, prints it, and opens + ``/auto-login/?token=...`` in the default browser. The token is checked in + constant time and consumed on first use, so the URL cannot be replayed by + another machine on the LAN (the demo server binds 0.0.0.0). Only enabled in MINIMAL_CONFIG (local demo) mode — 404 everywhere else. """ from django.conf import settings @@ -193,6 +200,12 @@ def auto_login_view(request): if not getattr(settings, "MINIMAL_CONFIG", False): raise Http404 + token = request.GET.get("token", "") + expected = os.environ.get("SIMPLEAUDIT_AUTO_LOGIN_TOKEN", "") + if not expected or not hmac.compare_digest(token, expected): + raise Http404 + # Single use: clear it so the URL stops working after this request. + os.environ.pop("SIMPLEAUDIT_AUTO_LOGIN_TOKEN", None) username = os.environ.get("BOOTSTRAP_USERNAME", "studio") user = User.objects.filter(username=username).first() if user is None: @@ -286,16 +299,12 @@ def post(self, request): def _grant_default_project(user): """Give first-time WorkOS users membership in the 'Default' project (viewer). - The 'Default' workspace is reserved for this purpose — it is created during - platform bootstrap and serves as the shared landing space for new users. + Thin wrapper over the shared ``grant_default_project`` service so every + user-creation path lands new users in the same shared landing space. """ - from accounts.models import Project, ProjectMembership + from accounts.services import grant_default_project - project = Project.objects.filter(slug="default").first() - if project: - ProjectMembership.objects.get_or_create( - project=project, user=user, defaults={"role": ProjectMembership.Role.VIEWER} - ) + grant_default_project(user) # ─── Dashboard ─────────────────────────────────────────────────────────────── @@ -1687,6 +1696,7 @@ class ConnectionsView(ProjectMixin, TemplateView): template_name = "connections.html" def get_context_data(self, **kw): + from model_registry import otlp_config from model_registry.models import ModelConnection from model_registry.services import ( PROVIDER_PRESETS, @@ -1736,6 +1746,12 @@ def get_context_data(self, **kw): connections=connections, conn_data=conn_data, model_total=sum(len(c.model_list) for c in connections), + # The per-model "open in chat" icon only makes sense when the chat + # module is on; otherwise /chat/ 404s. + chat_enabled=chat_enabled(), + # The OTLP credential button only makes sense while the OTLP + # listener is on; otherwise the issue/rotate endpoints 404. + otlp_enabled=otlp_config.ENABLED, provider_presets=PROVIDER_PRESETS, # Each provider once: presets share some (OpenAI and "Custom" are # both openai), and a connection's own provider must stay pickable. @@ -1928,6 +1944,105 @@ def post(self, request): return JsonResponse({"ok": True, "count": len(ids), "models": ids[:6]}) +class OTLPCredentialCreateView(ProjectMixin, View): + """Generate an OTLP credential for a connection. + + Returns the one-time secret plus the exact env vars to paste into the + target (Basic Auth or Bearer Token). Admin-only. + """ + + def post(self, request): + from model_registry import otlp_config + if not otlp_config.ENABLED: + return JsonResponse({"ok": False, "error": "OTLP is disabled on this deployment."}, status=404) + from model_registry import otlp_services as otlp + from model_registry.models import ModelConnection + from model_registry.otlp_views import _endpoint_url, _require_admin + + blocked = _require_write_access(request) + if blocked: + return JsonResponse({"ok": False, "error": "You don't have permission to change connections."}, status=403) + post = request.POST + conn = ModelConnection.objects.filter(pk=_int(post.get("connection_id")), project=request.project).first() + if conn is None: + return JsonResponse({"ok": False, "error": "Connection not found in this workspace."}, status=404) + try: + _require_admin(request, request.project) + except Exception as e: # noqa: BLE001 - surface the permission error + return JsonResponse({"ok": False, "error": str(e)}, status=403) + auth_mode = (post.get("auth_mode") or "basic").strip().lower() + if auth_mode not in ("none", "basic", "bearer"): + return JsonResponse({"ok": False, "error": "auth_mode must be 'none', 'basic', or 'bearer'."}, status=400) + try: + new_cred = otlp.create_credential(project=request.project, connection=conn, auth_mode=auth_mode, user=request.user) + except ValueError as e: + return JsonResponse({"ok": False, "error": str(e)}, status=400) + endpoint = _endpoint_url(request, origin=(post.get("origin") or "").strip()) + payload = { + "ok": True, + "credential": { + "id": new_cred.credential.id, + "auth_mode": new_cred.credential.auth_mode, + "username": new_cred.credential.username, + "target_id": new_cred.credential.target_id, + }, + "secret": new_cred.secret, + "endpoint": endpoint, + } + if new_cred.credential.auth_mode == "basic": + payload["env_vars"] = dict(otlp.otlp_env_vars(endpoint=endpoint, username=new_cred.credential.username, password=new_cred.secret)) + elif new_cred.credential.auth_mode == "bearer": + payload["env_vars"] = dict(otlp.otlp_bearer_env_vars(endpoint=endpoint, token=new_cred.secret)) + else: # none + payload["env_vars"] = dict(otlp.otlp_open_env_vars(endpoint=endpoint)) + return JsonResponse(payload) + + +class OTLPCredentialRotateView(ProjectMixin, View): + """Re-issue an OTLP credential's secret (old one stops working). Admin-only.""" + + def post(self, request): + from model_registry import otlp_config + if not otlp_config.ENABLED: + return JsonResponse({"ok": False, "error": "OTLP is disabled on this deployment."}, status=404) + from model_registry import otlp_services as otlp + from model_registry.models import OTLPCredential + from model_registry.otlp_views import _endpoint_url, _require_admin + + blocked = _require_write_access(request) + if blocked: + return JsonResponse({"ok": False, "error": "You don't have permission to change connections."}, status=403) + post = request.POST + cred = OTLPCredential.objects.filter(pk=_int(post.get("credential_id")), project=request.project).first() + if cred is None: + return JsonResponse({"ok": False, "error": "Credential not found in this workspace."}, status=404) + try: + _require_admin(request, request.project) + except Exception as e: # noqa: BLE001 + return JsonResponse({"ok": False, "error": str(e)}, status=403) + try: + new_cred = otlp.rotate_credential(cred) + except ValueError as e: + return JsonResponse({"ok": False, "error": str(e)}, status=400) + endpoint = _endpoint_url(request, origin=(post.get("origin") or "").strip()) + payload = { + "ok": True, + "credential": { + "id": new_cred.credential.id, + "auth_mode": new_cred.credential.auth_mode, + "username": new_cred.credential.username, + "target_id": new_cred.credential.target_id, + }, + "secret": new_cred.secret, + "endpoint": endpoint, + } + if new_cred.credential.auth_mode == "basic": + payload["env_vars"] = dict(otlp.otlp_env_vars(endpoint=endpoint, username=new_cred.credential.username, password=new_cred.secret)) + else: # bearer + payload["env_vars"] = dict(otlp.otlp_bearer_env_vars(endpoint=endpoint, token=new_cred.secret)) + return JsonResponse(payload) + + class DiscoverModelsView(ProjectMixin, View): """List the models a connection's server offers (GET {base_url}/models).""" diff --git a/infra/worker.py b/infra/worker.py index 925831c4..22eb21e0 100644 --- a/infra/worker.py +++ b/infra/worker.py @@ -364,6 +364,10 @@ def _scenario_execute_impl(workflow_input: ScenarioInput, ctx: Context) -> dict: gen_params = run.generation_parameters_snapshot or {} n_reps = int(gen_params.get("n_repetitions") or 1) max_turns = int(gen_params.get("max_turns") or 5) + # Trace acquisition config (Promptfoo parity). Empty = no tracing. Only the + # single-rep path forwards it: the engine's multi-rep path does not yet + # propagate trace correlation (see infra.engine.run_scenario_repeated). + trace_config = run.trace_config or None # Granular stage detail for the frontend: which phase of the scenario is # starting (target execution begins with the auditor generating a probe). @@ -532,6 +536,7 @@ def _stop_event_flusher() -> None: file_uri=revision.file_uri, category=item.scenario.category or "", metadata=revision.metadata or {}, + trace_config=trace_config, ) # Use aggregated severity for the run-level counter severity = result_payload.get("aggregated_severity", "") @@ -551,6 +556,7 @@ def _stop_event_flusher() -> None: file_uri=revision.file_uri, category=item.scenario.category or "", metadata=revision.metadata or {}, + trace_config=trace_config, ) severity = result_payload.get("severity", "") # Stop the live flusher: drains any remaining turn events and joins the diff --git a/model_registry/migrations/0004_otlpcredential.py b/model_registry/migrations/0004_otlpcredential.py new file mode 100644 index 00000000..b699ae94 --- /dev/null +++ b/model_registry/migrations/0004_otlpcredential.py @@ -0,0 +1,39 @@ +import django.db.models.deletion +from django.conf import settings +from django.db import migrations, models + + +class Migration(migrations.Migration): + + dependencies = [ + ("model_registry", "0003_modelconnection_shared_with_and_more"), + migrations.swappable_dependency(settings.AUTH_USER_MODEL), + ] + + operations = [ + migrations.CreateModel( + name="OTLPCredential", + fields=[ + ("id", models.BigAutoField(auto_created=True, primary_key=True, serialize=False, verbose_name="ID")), + ("auth_mode", models.CharField(choices=[("basic", "Basic Auth"), ("bearer", "Bearer Token")], default="basic", max_length=10)), + ("username", models.CharField(blank=True, default="", max_length=250)), + ("secret_hash", models.BinaryField()), + ("salt", models.BinaryField()), + ("target_id", models.CharField(max_length=250)), + ("enabled", models.BooleanField(default=True)), + ("created_at", models.DateTimeField(auto_now_add=True)), + ("updated_at", models.DateTimeField(auto_now=True)), + ("connection", models.ForeignKey(on_delete=django.db.models.deletion.CASCADE, related_name="otlp_credentials", to="model_registry.modelconnection")), + ("created_by", models.ForeignKey(blank=True, null=True, on_delete=django.db.models.deletion.SET_NULL, to=settings.AUTH_USER_MODEL)), + ("project", models.ForeignKey(on_delete=django.db.models.deletion.CASCADE, related_name="otlp_credentials", to="accounts.project")), + ], + options={ + "db_table": "core_otlp_credential", + "ordering": ["project__name", "target_id"], + }, + ), + migrations.AddConstraint( + model_name="otlpcredential", + constraint=models.UniqueConstraint(fields=["project", "target_id"], name="unique_target_id_per_project"), + ), + ] diff --git a/model_registry/migrations/0005_alter_otlpcredential_auth_mode_and_more.py b/model_registry/migrations/0005_alter_otlpcredential_auth_mode_and_more.py new file mode 100644 index 00000000..f0003d1e --- /dev/null +++ b/model_registry/migrations/0005_alter_otlpcredential_auth_mode_and_more.py @@ -0,0 +1,35 @@ +# Generated by Django 5.2.4 on 2026-10-01 20:11 + +from django.db import migrations, models + + +class Migration(migrations.Migration): + dependencies = [ + ("model_registry", "0004_otlpcredential"), + ] + + operations = [ + migrations.AlterField( + model_name="otlpcredential", + name="auth_mode", + field=models.CharField( + choices=[ + ("none", "None (open)"), + ("basic", "Basic Auth"), + ("bearer", "Bearer Token"), + ], + default="basic", + max_length=10, + ), + ), + migrations.AlterField( + model_name="otlpcredential", + name="salt", + field=models.BinaryField(blank=True, null=True), + ), + migrations.AlterField( + model_name="otlpcredential", + name="secret_hash", + field=models.BinaryField(blank=True, null=True), + ), + ] diff --git a/model_registry/migrations/0006_otlpcredential_token_prefix.py b/model_registry/migrations/0006_otlpcredential_token_prefix.py new file mode 100644 index 00000000..ef6057e9 --- /dev/null +++ b/model_registry/migrations/0006_otlpcredential_token_prefix.py @@ -0,0 +1,19 @@ +# Generated by Django 5.2.4 on 2026-10-01 21:18 + +from django.db import migrations, models + + +class Migration(migrations.Migration): + dependencies = [ + ("model_registry", "0005_alter_otlpcredential_auth_mode_and_more"), + ] + + operations = [ + migrations.AddField( + model_name="otlpcredential", + name="token_prefix", + field=models.CharField( + blank=True, db_index=True, default="", max_length=32 + ), + ), + ] diff --git a/model_registry/models.py b/model_registry/models.py index d7744a97..a1fa391e 100644 --- a/model_registry/models.py +++ b/model_registry/models.py @@ -98,3 +98,54 @@ def has_key(self) -> bool: """Whether the parent connection has an API key configured.""" return self.connection.has_key + +class OTLPCredential(models.Model): + """A credential that lets an external target push OTLP traces to Studio. + + One credential per target. The target is configured with either a Basic + username+password (standard ``OTEL_BASIC_AUTH_*``) or a Bearer + token (generic OTel exporters). The receiver authenticates the request and + tags every ingested span with ``target_id`` so spans are attributable. + + Only a salted hash of the secret is stored — the plaintext password/token + is shown once at creation and cannot be read back. + """ + + class AuthMode(models.TextChoices): + NONE = "none", "None (open)" + BASIC = "basic", "Basic Auth" + BEARER = "bearer", "Bearer Token" + + project = models.ForeignKey("accounts.Project", on_delete=models.CASCADE, related_name="otlp_credentials") + connection = models.ForeignKey(ModelConnection, on_delete=models.CASCADE, related_name="otlp_credentials") + auth_mode = models.CharField(max_length=10, choices=AuthMode.choices, default=AuthMode.BASIC) + # For basic: the username the target sends. For bearer: a stable label. + username = models.CharField(max_length=250, blank=True, default="") + # Salted SHA-256 of the secret (password for basic, token for bearer). + secret_hash = models.BinaryField(null=True, blank=True) + salt = models.BinaryField(null=True, blank=True) + # Non-secret lookup prefix of a bearer token (e.g. "sa_otlp_abc123"). Lets + # verify_bearer do an indexed lookup instead of hashing against every + # credential. Never used for authentication — only the salted hash is. + token_prefix = models.CharField(max_length=32, blank=True, default="", db_index=True) + # Stable identifier stamped onto every span this credential authenticates. + target_id = models.CharField(max_length=250) + enabled = models.BooleanField(default=True) + created_by = models.ForeignKey(settings.AUTH_USER_MODEL, on_delete=models.SET_NULL, null=True, blank=True) + created_at = models.DateTimeField(auto_now_add=True) + updated_at = models.DateTimeField(auto_now=True) + + class Meta: + db_table = "core_otlp_credential" + constraints = [ + models.UniqueConstraint(fields=["project", "target_id"], name="unique_target_id_per_project"), + ] + ordering = ["project__name", "target_id"] + + def __str__(self) -> str: + return f"OTLP {self.auth_mode}:{self.target_id}" + + @property + def display_name(self) -> str: + return self.username or self.target_id + diff --git a/model_registry/otlp_config.py b/model_registry/otlp_config.py new file mode 100644 index 00000000..8c2cdddf --- /dev/null +++ b/model_registry/otlp_config.py @@ -0,0 +1,30 @@ +"""Whether Studio's OTLP listener is switched on. + +The OTLP listener (``POST /otlp/v1/traces``) plus its credential-management +API and UI are a self-contained Studio feature: they let an external target +push OpenTelemetry spans in and let an admin issue the credentials that gate +them. They have no bearing on the audit engine itself, so a deployment that +doesn't want them can switch the whole thing off. + +**On by default** to preserve the existing behavior: the listener has always +been wired unconditionally, and existing chat-OTLP setups depend on it. Set +``SIMPLEAUDIT_OTLP`` to any of the "off" spellings below to 404 the +``/otlp/*`` routes and the credential UI. (Unlike the chat module — which is +*off* until enabled — this flag is *on* until explicitly disabled, so leaving +the variable unset keeps today's behavior.) +""" +from __future__ import annotations + +import os + +#: Spellings of "OTLP off". An unset/empty variable is NOT in this set, so the +#: listener stays on by default. +DISABLED_VALUES = frozenset({"off", "disabled", "disable", "false", "no", "none", "0"}) + + +def is_disabled(value: str | None) -> bool: + return (value or "").strip().lower() in DISABLED_VALUES + + +#: The OTLP listener is on unless explicitly switched off. +ENABLED = not is_disabled(os.environ.get("SIMPLEAUDIT_OTLP")) diff --git a/model_registry/otlp_services.py b/model_registry/otlp_services.py new file mode 100644 index 00000000..9151ca6e --- /dev/null +++ b/model_registry/otlp_services.py @@ -0,0 +1,200 @@ +"""OTLP credential services for Studio. + +Generates and verifies the credentials an external target uses to push OTLP +traces to Studio. The credential *primitive* — salted SHA-256 hashing, +constant-time verification, header parsing, and secret generation — lives in +the core (``simpleaudit.tracing.auth``) so the library and Studio share one +definition. This module adds the Studio persistence layer on top: the +``OTLPCredential`` table, per-connection lookup, and the "none" fallback. + +Two auth modes: + - ``basic`` — username + password (OpenWebUI's native OTEL_BASIC_AUTH_*). + - ``bearer`` — a single token (generic OTel exporters via headers). + +Only a salted hash of the secret is stored; the plaintext is returned once at +creation so the user can copy it into the target's environment. +""" +from __future__ import annotations + +import secrets +from dataclasses import dataclass + +from django.db import transaction + +# The credential primitive (salted-hash + constant-time verify + header +# parsing + secret generation) lives in the core so the library and Studio +# share one definition. Studio adds the persistence layer on top: the +# OTLPCredential table, per-connection lookup, and the "none" fallback. +from simpleaudit.tracing.auth import ( + generate_password, + generate_token, + hash_secret, + new_salt, + token_lookup_prefix, + verify_secret, +) + +from model_registry.models import OTLPCredential + + +def _make_target_id(connection) -> str: + """A stable, human-readable target id derived from the connection.""" + base = "".join(c for c in connection.name.lower() if c.isalnum())[:24] or "target" + return f"{base}_{secrets.token_hex(4)}" + + +@dataclass +class NewCredential: + """A freshly created credential plus its one-time plaintext secret.""" + + credential: OTLPCredential + secret: str # password (basic) or token (bearer) + + +@transaction.atomic +def create_credential(*, project, connection, auth_mode: str, user=None) -> NewCredential: + """Create a credential for ``connection`` and return it with its secret. + + ``auth_mode`` is ``"none"``, ``"basic"``, or ``"bearer"``. For ``none`` no + secret is generated (the endpoint is open); for the other two the plaintext + secret is generated here and returned once, with only its salted hash + persisted. + """ + valid = (OTLPCredential.AuthMode.NONE, OTLPCredential.AuthMode.BASIC, OTLPCredential.AuthMode.BEARER) + if auth_mode not in valid: + raise ValueError(f"Unknown OTLP auth mode: {auth_mode!r}") + + # Reuse the existing target_id if a credential already exists for this + # connection, so the upsert updates in place rather than creating a new + # target identity. + existing = OTLPCredential.objects.filter(project=project, connection=connection).first() + target_id = existing.target_id if existing else _make_target_id(connection) + token_prefix = "" + if auth_mode == OTLPCredential.AuthMode.NONE: + username = "" + secret = "" + salt = None + secret_hash = None + elif auth_mode == OTLPCredential.AuthMode.BASIC: + username = f"sa_{target_id}" + secret = generate_password() + salt = new_salt() + secret_hash = hash_secret(secret, salt) + else: # bearer + username = "" + secret = generate_token() + salt = new_salt() + secret_hash = hash_secret(secret, salt) + token_prefix = token_lookup_prefix(secret) + + # Upsert: if a credential already exists for this connection, update it + # in place (change auth mode, regenerate secret) instead of failing on + # the unique constraint. + cred, _created = OTLPCredential.objects.update_or_create( + project=project, + target_id=target_id, + defaults={ + "connection": connection, + "auth_mode": auth_mode, + "username": username, + "secret_hash": secret_hash, + "salt": salt, + "token_prefix": token_prefix if auth_mode == OTLPCredential.AuthMode.BEARER else "", + "enabled": True, + "created_by": user, + }, + ) + return NewCredential(credential=cred, secret=secret) + + +@transaction.atomic +def rotate_credential(cred: OTLPCredential) -> NewCredential: + """Re-issue the secret for an existing credential (same target_id). + + The old secret stops working immediately. Returns the credential plus the + new one-time plaintext secret. Only meaningful for ``basic``/``bearer`` — + a ``none`` credential has no secret to rotate. + """ + if cred.auth_mode == OTLPCredential.AuthMode.NONE: + raise ValueError("A 'none' credential has no secret to rotate.") + + if cred.auth_mode == OTLPCredential.AuthMode.BASIC: + secret = generate_password() + else: # bearer + secret = generate_token() + + salt = new_salt() + cred.secret_hash = hash_secret(secret, salt) + cred.salt = salt + cred.token_prefix = token_lookup_prefix(secret) + cred.enabled = True # rotation re-enables a revoked credential + cred.save(update_fields=["secret_hash", "salt", "token_prefix", "enabled", "updated_at"]) + return NewCredential(credential=cred, secret=secret) + + +def verify_basic(username: str, password: str) -> OTLPCredential | None: + """Return the enabled credential matching ``username``/``password``, else None.""" + cred = OTLPCredential.objects.filter( + username=username, auth_mode=OTLPCredential.AuthMode.BASIC, enabled=True + ).first() + if cred is None: + return None + if verify_secret(password, cred.salt, cred.secret_hash): + return cred + return None + + +def verify_bearer(token: str) -> OTLPCredential | None: + """Return the enabled bearer credential matching ``token``, else None. + + The token's short lookup prefix narrows the candidate set (an indexed + query), then the salted hash is verified in constant time. Falls back to + scanning all bearer credentials when the prefix is unset (e.g. rows + created before the ``token_prefix`` field existed) so verification never + regresses. + """ + prefix = token_lookup_prefix(token) + qs = OTLPCredential.objects.filter( + auth_mode=OTLPCredential.AuthMode.BEARER, enabled=True + ).only("salt", "secret_hash", "target_id", "id", "token_prefix") + # Prefer the indexed prefix match; if none is stored, scan the (small) set. + candidates = qs.filter(token_prefix=prefix) + if not candidates.exists(): + candidates = qs + for cred in candidates: + if verify_secret(token, cred.salt, cred.secret_hash): + return cred + return None + + +def otlp_env_vars(*, endpoint: str, username: str, password: str) -> list[tuple[str, str]]: + """Env vars for a target that authenticates with Basic Auth. + + These map to the standard ``OTEL_BASIC_AUTH_USERNAME`` / ``PASSWORD`` + variables that OpenTelemetry exporters (e.g. OpenWebUI's) convert into an + ``Authorization: Basic ...`` header. + """ + return [ + ("OTEL_OTLP_SPAN_EXPORTER", "http"), + ("OTEL_EXPORTER_OTLP_ENDPOINT", endpoint), + ("OTEL_EXPORTER_OTLP_PROTOCOL", "http/json"), + ("OTEL_BASIC_AUTH_USERNAME", username), + ("OTEL_BASIC_AUTH_PASSWORD", password), + ] + + +def otlp_bearer_env_vars(*, endpoint: str, token: str) -> list[tuple[str, str]]: + """Env vars for a generic OTel exporter using a Bearer token.""" + return [ + ("OTEL_EXPORTER_OTLP_ENDPOINT", endpoint), + ("OTEL_EXPORTER_OTLP_PROTOCOL", "http/json"), + ("OTEL_EXPORTER_OTLP_HEADERS", f"Authorization=Bearer {token}"), + ] + + +def otlp_open_env_vars(*, endpoint: str) -> list[tuple[str, str]]: + """Env vars for an unauthenticated (``none``) target — just the endpoint.""" + return [ + ("OTEL_EXPORTER_OTLP_ENDPOINT", endpoint), + ("OTEL_EXPORTER_OTLP_PROTOCOL", "http/json"), + ] diff --git a/model_registry/otlp_views.py b/model_registry/otlp_views.py new file mode 100644 index 00000000..d7519f6a --- /dev/null +++ b/model_registry/otlp_views.py @@ -0,0 +1,248 @@ +"""OTLP ingestion + credential management for Studio. + +Two kinds of endpoints live here: + +1. **Ingestion** — ``POST /otlp/v1/traces``. Machine-to-machine: an external + target (OpenWebUI, an agent, a framework) pushes OTLP/HTTP-JSON spans. It is + authenticated by the ``Authorization`` header (Basic or Bearer) against the + workspace's ``OTLPCredential`` rows, NOT by a Studio session. On success the + spans are tagged with the credential's ``target_id`` and stored. + +2. **Credential management** — session-authenticated DRF views under + ``/api/otlp/credentials/`` that let a workspace admin create, list, and + revoke credentials (and fetch the copy-paste env vars for a target). +""" +from __future__ import annotations + +from django.http import JsonResponse +from django.views.decorators.csrf import csrf_exempt +from rest_framework.decorators import api_view, permission_classes +from rest_framework.permissions import IsAuthenticated +from rest_framework.response import Response +from simpleaudit.tracing.auth import parse_basic_header, parse_bearer_header +from simpleaudit.tracing.store import SpanStore + +from infra.exceptions import StableAPIError +from model_registry import otlp_services as otlp +from model_registry.models import ModelConnection, OTLPCredential + +# In-memory span store for the shared receiver, keyed by target_id. Each bucket +# is an engine ``SpanStore`` (the same normalized schema the run path and the +# judge's evidence selection use), so spans ingested here are directly +# consumable by ``select_spans``. (A persistent backend can replace this +# without changing the endpoint contract.) +_SPAN_STORE: dict[str, SpanStore] = {} + + +def _store_for_target(target_id: str) -> SpanStore: + store = _SPAN_STORE.get(target_id) + if store is None: + store = SpanStore() + _SPAN_STORE[target_id] = store + return store + + +def get_spans_for_target(target_id: str) -> list[dict]: + """All normalized spans ingested for ``target_id`` (for a run to fetch evidence).""" + store = _SPAN_STORE.get(target_id) + return store.all() if store else [] + + +def clear_target_spans(target_id: str) -> None: + _SPAN_STORE.pop(target_id, None) + + +# ─── Ingestion (machine-to-machine) ───────────────────────────────────────── + + +@csrf_exempt +def otlp_traces(request): + """Accept an OTLP/HTTP-JSON trace export from an authenticated target. + + Auth: ``Authorization: Basic base64(user:pass)`` or ``Bearer ``. + Returns the OTLP ack (200) on success, 401 on auth failure, 405 on a + non-POST. Spans are tagged with the credential's ``target_id``. + """ + if request.method != "POST": + return JsonResponse({"error": {"code": "method_not_allowed", "message": "Use POST."}}, status=405) + + authorization = request.headers.get("Authorization") + cred = None + if authorization and authorization.strip().lower().startswith("basic"): + parsed = parse_basic_header(authorization) + if parsed: + cred = otlp.verify_basic(*parsed) + elif authorization and authorization.strip().lower().startswith("bearer"): + token = parse_bearer_header(authorization) + if token: + cred = otlp.verify_bearer(token) + + # No Authorization header (or an unrecognized scheme): fall back to an + # enabled "none" credential for this workspace, if any. This is the + # unauthenticated path for targets that can't send credentials. + if cred is None: + cred = OTLPCredential.objects.filter( + auth_mode=OTLPCredential.AuthMode.NONE, enabled=True + ).first() + + if cred is None: + return JsonResponse( + {"error": {"code": "unauthorized", "message": "Invalid or missing OTLP credentials."}}, + status=401, + ) + + try: + from simpleaudit.tracing.otlp import parse_otlp_json + + # Parse the export, tag each span with the authenticated target so it's + # attributable, then normalize into the shared SpanStore (the same + # schema the run path and the judge's evidence selection use). + raw_spans = parse_otlp_json(request.body) + for span in raw_spans: + attrs = span.get("attributes") + if attrs is None: + span["attributes"] = attrs = {} + attrs.setdefault("simpleaudit.target_id", cred.target_id) + _store_for_target(cred.target_id).add_many(raw_spans) + return JsonResponse( + {"partialSuccess": {"rejectedSpans": 0}, "authenticated": True}, status=200 + ) + except Exception: # noqa: BLE001 - never fail the export; ack a rejection + return JsonResponse({"partialSuccess": {"rejectedSpans": 1}}, status=200) + + +# ─── Credential management (session-authenticated) ────────────────────────── + + +def _active_project(request): + project = getattr(request, "project", None) + if project is None: + raise StableAPIError(detail="No active workspace.", code="no_workspace", http_status=400) + return project + + +def _require_admin(request, project) -> None: + from accounts.services import _require_workspace_admin + + _require_workspace_admin(request.user, project) + + +def _endpoint_url(request, origin: str | None = None) -> str: + """The absolute OTLP traces URL a target should export to. + + Prefers the browser's own ``origin`` (sent by the UI) so the URL reflects + what the user actually sees — the real domain, not a reverse-proxy or + 127.0.0.1 artifact. Falls back to the request's Host header. + """ + origin = (origin or "").strip() + if origin.startswith(("http://", "https://")): + return f"{origin.rstrip('/')}/otlp/v1/traces" + host = request.headers.get("Host") or request.get_host() + scheme = "https" if request.is_secure() else "http" + return f"{scheme}://{host}/otlp/v1/traces" + + +def _serialize_cred(cred: OTLPCredential, request) -> dict: + return { + "id": cred.id, + "auth_mode": cred.auth_mode, + "username": cred.username, + "target_id": cred.target_id, + "display_name": cred.display_name, + "enabled": cred.enabled, + "created_at": cred.created_at.isoformat() if cred.created_at else None, + } + + +@api_view(["GET"]) +@permission_classes([IsAuthenticated]) +def list_credentials(request): + """List this workspace's OTLP credentials (no secrets).""" + project = _active_project(request) + creds = OTLPCredential.objects.filter(project=project).select_related("connection") + return Response({"credentials": [_serialize_cred(c, request) for c in creds]}) + + +@api_view(["POST"]) +@permission_classes([IsAuthenticated]) +def create_credential(request): + """Create a credential for a connection. Returns the one-time secret. + + Body: ``{"connection_id": int, "auth_mode": "basic"|"bearer"}``. + The response includes the plaintext password/token exactly once. + """ + project = _active_project(request) + _require_admin(request, project) + data = request.data or {} + conn_id = data.get("connection_id") + auth_mode = (data.get("auth_mode") or "basic").strip().lower() + if not conn_id or not str(conn_id).isdigit(): + raise StableAPIError(detail="connection_id is required.", code="invalid_connection_id") + conn = ModelConnection.objects.filter(pk=int(conn_id), project=project).first() + if conn is None: + raise StableAPIError(detail="Connection not found in this workspace.", code="conn_not_found", http_status=404) + + new_cred = otlp.create_credential(project=project, connection=conn, auth_mode=auth_mode, user=request.user) + endpoint = _endpoint_url(request, origin=(request.data.get("origin") or "").strip()) + payload = { + "credential": _serialize_cred(new_cred.credential, request), + "secret": new_cred.secret, + "endpoint": endpoint, + } + if new_cred.credential.auth_mode == OTLPCredential.AuthMode.BASIC: + payload["env_vars"] = dict( + otlp.otlp_env_vars(endpoint=endpoint, username=new_cred.credential.username, password=new_cred.secret) + ) + elif new_cred.credential.auth_mode == OTLPCredential.AuthMode.BEARER: + payload["env_vars"] = dict(otlp.otlp_bearer_env_vars(endpoint=endpoint, token=new_cred.secret)) + else: # none — no credentials to paste; just the endpoint + payload["env_vars"] = dict( + otlp.otlp_open_env_vars(endpoint=endpoint) + ) + return Response(payload) + + +@api_view(["POST"]) +@permission_classes([IsAuthenticated]) +def rotate_credential(request, cred_id): + """Re-issue a credential's secret. The old secret stops working. + + Returns the new one-time secret plus the env vars to paste into the target. + """ + project = _active_project(request) + _require_admin(request, project) + cred = OTLPCredential.objects.filter(pk=cred_id, project=project).first() + if cred is None: + raise StableAPIError(detail="Credential not found.", code="cred_not_found", http_status=404) + try: + new_cred = otlp.rotate_credential(cred) + except ValueError as e: + raise StableAPIError(detail=str(e), code="cannot_rotate", http_status=400) + endpoint = _endpoint_url(request, origin=(request.data.get("origin") or "").strip()) + payload = { + "credential": _serialize_cred(new_cred.credential, request), + "secret": new_cred.secret, + "endpoint": endpoint, + } + if new_cred.credential.auth_mode == OTLPCredential.AuthMode.BASIC: + payload["env_vars"] = dict( + otlp.otlp_env_vars(endpoint=endpoint, username=new_cred.credential.username, password=new_cred.secret) + ) + else: # bearer + payload["env_vars"] = dict(otlp.otlp_bearer_env_vars(endpoint=endpoint, token=new_cred.secret)) + return Response(payload) + + +@api_view(["POST"]) +@permission_classes([IsAuthenticated]) +def revoke_credential(request, cred_id): + """Revoke (disable) a credential so it can no longer push traces.""" + project = _active_project(request) + _require_admin(request, project) + cred = OTLPCredential.objects.filter(pk=cred_id, project=project).first() + if cred is None: + raise StableAPIError(detail="Credential not found.", code="cred_not_found", http_status=404) + cred.enabled = False + cred.save(update_fields=["enabled", "updated_at"]) + otlp.clear_target_spans(cred.target_id) + return Response({"ok": True, "credential": _serialize_cred(cred, request)}) diff --git a/model_registry/tests/__init__.py b/model_registry/tests/__init__.py new file mode 100644 index 00000000..e69de29b diff --git a/model_registry/tests/test_otlp_receiver.py b/model_registry/tests/test_otlp_receiver.py new file mode 100644 index 00000000..d89c687d --- /dev/null +++ b/model_registry/tests/test_otlp_receiver.py @@ -0,0 +1,295 @@ +"""Tests for the shared OTLP receiver: credential services + ingestion endpoint. + +Covers the two auth modes (basic / bearer), header parsing, span tagging, +auth failures, method handling, and revocation. +""" +import base64 +import json + +from django.test import Client, TestCase +from simpleaudit.tracing.auth import parse_basic_header, parse_bearer_header + +from infra.tests.factories import ( + MembershipFactory, + ModelConnectionFactory, + ProjectFactory, + UserFactory, +) +from model_registry import otlp_services as otlp +from model_registry import otlp_views +from model_registry.models import OTLPCredential + + +def _basic_header(username, password): + return "Basic " + base64.b64encode(f"{username}:{password}".encode()).decode() + + +def _otlp_body(): + return json.dumps( + {"resourceSpans": [{"scopeSpans": [{"spans": [{"traceId": "ab" * 16, "spanId": "cd" * 8, "name": "turn"}]}]}]} + ) + + +class OTLPHeaderParsingTest(TestCase): + def test_parse_basic(self): + self.assertEqual(parse_basic_header(_basic_header("u", "p")), ("u", "p")) + + def test_parse_basic_with_colon_in_password(self): + self.assertEqual(parse_basic_header(_basic_header("u", "a:b:c")), ("u", "a:b:c")) + + def test_parse_bearer(self): + self.assertEqual(parse_bearer_header("Bearer tok123"), "tok123") + + def test_scheme_mismatch_returns_none(self): + self.assertIsNone(parse_basic_header("Bearer x")) + self.assertIsNone(parse_bearer_header("Basic x")) + + def test_missing_header(self): + self.assertIsNone(parse_basic_header(None)) + self.assertIsNone(parse_bearer_header("")) + + def test_invalid_base64(self): + self.assertIsNone(parse_basic_header("Basic !!!not-base64!!!")) + + +class OTLPCredentialServiceTest(TestCase): + def setUp(self): + self.user = UserFactory() + self.project = ProjectFactory() + self.conn = ModelConnectionFactory(project=self.project, name="owui-a") + + def test_create_basic_returns_secret_and_stores_hash(self): + nc = otlp.create_credential(project=self.project, connection=self.conn, auth_mode="basic", user=self.user) + self.assertTrue(nc.secret) + self.assertEqual(nc.credential.auth_mode, OTLPCredential.AuthMode.BASIC) + self.assertTrue(nc.credential.username.startswith("sa_")) + # The plaintext is not stored; only a hash. + self.assertNotEqual(nc.credential.secret_hash, nc.secret.encode()) + # Round-trip verification succeeds. + self.assertIsNotNone(otlp.verify_basic(nc.credential.username, nc.secret)) + + def test_create_bearer_returns_token(self): + nb = otlp.create_credential(project=self.project, connection=self.conn, auth_mode="bearer", user=self.user) + self.assertTrue(nb.secret.startswith("sa_otlp_")) + self.assertEqual(nb.credential.username, "") + self.assertIsNotNone(otlp.verify_bearer(nb.secret)) + + def test_verify_wrong_password_returns_none(self): + nc = otlp.create_credential(project=self.project, connection=self.conn, auth_mode="basic", user=self.user) + self.assertIsNone(otlp.verify_basic(nc.credential.username, "wrong")) + + def test_verify_wrong_bearer_returns_none(self): + self.assertIsNone(otlp.verify_bearer("sa_otlp_not_a_real_token")) + + def test_bearer_stores_lookup_prefix(self): + nb = otlp.create_credential(project=self.project, connection=self.conn, auth_mode="bearer", user=self.user) + self.assertEqual(nb.credential.token_prefix, otlp.token_lookup_prefix(nb.secret)) + # A token sharing the prefix but differing in the hash body must not match. + fake = nb.secret[:16] + "0" * (len(nb.secret) - 16) + self.assertNotEqual(fake, nb.secret) + self.assertIsNone(otlp.verify_bearer(fake)) + + def test_bearer_legacy_row_without_prefix_still_verifies(self): + """Rows created before token_prefix existed (empty prefix) still verify via fallback.""" + nb = otlp.create_credential(project=self.project, connection=self.conn, auth_mode="bearer", user=self.user) + nb.credential.token_prefix = "" + nb.credential.save(update_fields=["token_prefix"]) + self.assertIsNotNone(otlp.verify_bearer(nb.secret)) + + def test_disabled_credential_not_verified(self): + nc = otlp.create_credential(project=self.project, connection=self.conn, auth_mode="basic", user=self.user) + nc.credential.enabled = False + nc.credential.save() + self.assertIsNone(otlp.verify_basic(nc.credential.username, nc.secret)) + + def test_invalid_mode_raises(self): + with self.assertRaises(ValueError): + otlp.create_credential(project=self.project, connection=self.conn, auth_mode="mtls") + + def test_create_none_has_no_secret(self): + nc = otlp.create_credential(project=self.project, connection=self.conn, auth_mode="none", user=self.user) + self.assertEqual(nc.credential.auth_mode, OTLPCredential.AuthMode.NONE) + self.assertEqual(nc.secret, "") + self.assertIsNone(nc.credential.secret_hash) + self.assertIsNone(nc.credential.salt) + + def test_rotate_basic(self): + """Rotation re-issues the secret; the old one stops working.""" + cred = otlp.create_credential(project=self.project, connection=self.conn, auth_mode="basic", user=self.user) + old = cred.secret + new = otlp.rotate_credential(cred.credential) + self.assertNotEqual(new.secret, old) + self.assertIsNone(otlp.verify_basic(cred.credential.username, old)) + self.assertIsNotNone(otlp.verify_basic(cred.credential.username, new.secret)) + + def test_rotate_bearer(self): + cred = otlp.create_credential(project=self.project, connection=self.conn, auth_mode="bearer", user=self.user) + old = cred.secret + new = otlp.rotate_credential(cred.credential) + self.assertNotEqual(new.secret, old) + self.assertIsNone(otlp.verify_bearer(old)) + self.assertIsNotNone(otlp.verify_bearer(new.secret)) + + def test_rotate_reenables_revoked(self): + cred = otlp.create_credential(project=self.project, connection=self.conn, auth_mode="basic", user=self.user) + cred.credential.enabled = False + cred.credential.save() + new = otlp.rotate_credential(cred.credential) + self.assertTrue(new.credential.enabled) + self.assertIsNotNone(otlp.verify_basic(new.credential.username, new.secret)) + + def test_rotate_none_raises(self): + cred = otlp.create_credential(project=self.project, connection=self.conn, auth_mode="none", user=self.user) + with self.assertRaises(ValueError): + otlp.rotate_credential(cred.credential) + + +class OTLPIngestionEndpointTest(TestCase): + def setUp(self): + self.user = UserFactory() + self.project = ProjectFactory() + MembershipFactory(user=self.user, project=self.project, role="admin") + self.conn = ModelConnectionFactory(project=self.project, name="owui-a") + self.client = Client(SERVER_NAME="localhost") + + def _cred(self, mode): + return otlp.create_credential(project=self.project, connection=self.conn, auth_mode=mode, user=self.user) + + def test_valid_basic_ingests_and_tags(self): + nc = self._cred("basic") + otlp_views.clear_target_spans(nc.credential.target_id) + resp = self.client.post( + "/otlp/v1/traces", data=_otlp_body(), content_type="application/json", + HTTP_AUTHORIZATION=_basic_header(nc.credential.username, nc.secret), + ) + self.assertEqual(resp.status_code, 200) + self.assertTrue(resp.json()["authenticated"]) + spans = otlp_views.get_spans_for_target(nc.credential.target_id) + self.assertEqual(len(spans), 1) + self.assertEqual(spans[0]["attributes"]["simpleaudit.target_id"], nc.credential.target_id) + + def test_valid_bearer_ingests(self): + nb = self._cred("bearer") + resp = self.client.post( + "/otlp/v1/traces", data=_otlp_body(), content_type="application/json", + HTTP_AUTHORIZATION=f"Bearer {nb.secret}", + ) + self.assertEqual(resp.status_code, 200) + self.assertEqual(len(otlp_views.get_spans_for_target(nb.credential.target_id)), 1) + + def test_wrong_password_401(self): + nc = self._cred("basic") + resp = self.client.post( + "/otlp/v1/traces", data=_otlp_body(), content_type="application/json", + HTTP_AUTHORIZATION=_basic_header(nc.credential.username, "wrong"), + ) + self.assertEqual(resp.status_code, 401) + + def test_no_auth_401(self): + resp = self.client.post("/otlp/v1/traces", data=_otlp_body(), content_type="application/json") + self.assertEqual(resp.status_code, 401) + + def test_get_rejected(self): + self.assertEqual(self.client.get("/otlp/v1/traces").status_code, 405) + + def test_revoke_blocks_subsequent_push(self): + nc = self._cred("basic") + header = _basic_header(nc.credential.username, nc.secret) + self.assertEqual( + self.client.post("/otlp/v1/traces", data=_otlp_body(), content_type="application/json", + HTTP_AUTHORIZATION=header).status_code, 200) + nc.credential.enabled = False + nc.credential.save() + self.assertEqual( + self.client.post("/otlp/v1/traces", data=_otlp_body(), content_type="application/json", + HTTP_AUTHORIZATION=header).status_code, 401) + + def test_malformed_body_ack_not_500(self): + nc = self._cred("basic") + resp = self.client.post( + "/otlp/v1/traces", data="{not json", content_type="application/json", + HTTP_AUTHORIZATION=_basic_header(nc.credential.username, nc.secret), + ) + self.assertEqual(resp.status_code, 200) + self.assertEqual(resp.json()["partialSuccess"]["rejectedSpans"], 1) + + def test_none_mode_accepts_unauthenticated(self): + nc = self._cred("none") + resp = self.client.post("/otlp/v1/traces", data=_otlp_body(), content_type="application/json") + self.assertEqual(resp.status_code, 200) + self.assertTrue(resp.json()["authenticated"]) + spans = otlp_views.get_spans_for_target(nc.credential.target_id) + self.assertEqual(len(spans), 1) + self.assertEqual(spans[0]["attributes"]["simpleaudit.target_id"], nc.credential.target_id) + + def test_none_mode_disabled_blocks_unauthenticated(self): + nc = self._cred("none") + nc.credential.enabled = False + nc.credential.save() + resp = self.client.post("/otlp/v1/traces", data=_otlp_body(), content_type="application/json") + self.assertEqual(resp.status_code, 401) + + + +class OTLPFlagTest(TestCase): + """The OTLP listener is on by default and can be switched off. + + The routes are wired at URLconf import time based on + ``model_registry.otlp_config.ENABLED``. When off, the ``/otlp/*`` and + ``/api/otlp/*`` paths are not registered, so a deployment that doesn't + want the listener exposes no OTLP surface. + """ + + def setUp(self): + self.user = UserFactory() + self.project = ProjectFactory() + MembershipFactory(user=self.user, project=self.project, role="admin") + self.client = Client(SERVER_NAME="localhost") + + def _reload_urlconf(self): + """Re-import config.urls so the otlp routes reflect the current flag.""" + import importlib + + from django.urls import clear_url_caches, set_urlconf + + import config.urls + + importlib.reload(config.urls) + set_urlconf("config.urls") + clear_url_caches() + + def _resolvable(self, name): + """Whether a named route currently resolves (fresh resolver).""" + from django.urls import NoReverseMatch, reverse + + try: + reverse(name) + return True + except NoReverseMatch: + return False + + def test_enabled_by_default(self): + from model_registry import otlp_config + + self.assertTrue(otlp_config.ENABLED) + + def test_routes_present_when_enabled(self): + from model_registry import otlp_config + + self.assertTrue(otlp_config.ENABLED) + self._reload_urlconf() + for name in ("otlp-traces", "otlp-credentials-list", "otlp-credentials-create"): + self.assertTrue(self._resolvable(name), f"{name} should resolve when enabled") + + def test_routes_absent_when_disabled(self): + from model_registry import otlp_config + + old = otlp_config.ENABLED + otlp_config.ENABLED = False + try: + self._reload_urlconf() + for name in ("otlp-traces", "otlp-credentials-list", "otlp-credentials-create"): + self.assertFalse(self._resolvable(name), f"{name} should not resolve when disabled") + finally: + otlp_config.ENABLED = old + self._reload_urlconf() diff --git a/pyproject.toml b/pyproject.toml index d168c2ba..eb37318e 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -4,7 +4,7 @@ build-backend = "hatchling.build" [project] name = "simpleaudit-studio" -version = "0.5.2" +version = "0.5.3" description = "Production AI model audit platform — one-liner demo mode via uvx" readme = "README.md" requires-python = ">=3.11" @@ -26,14 +26,33 @@ dependencies = [ # 0.2.1 adds the native on_turn per-phase progress callback. "simpleaudit==0.2.3", "cronsim==2.7", + "psutil>=7.2.2", ] [project.optional-dependencies] postgres = ["psycopg[binary]==3.2.13"] server = ["gunicorn==23.0.0"] -dev = ["factory_boy==3.3.2", "ruff"] +dev = [ + "factory_boy==3.3.2", + "ruff", + # Test runner stack. xdist is pinned <3.8 because 3.8+ pulls in `greenlet`, + # which does not build on Python 3.11 (the project's lower bound). + "pytest>=8,<9", + "pytest-django>=4.8,<5", + "pytest-xdist>=3.6,<3.8", + "pytest-testmon>=2,<3", + "pytest-env>=1,<2", +] e2e = ["playwright>=1.45"] +# Local-dev override: point the engine at the in-repo SimpleAudit checkout so +# the studio can use the tracing layer (simpleaudit.tracing) that is not yet in +# the published 0.2.3 wheel. The published build still resolves simpleaudit +# from PyPI (see the registry-dependency note above); uv.sources only affects +# local `uv` resolution. +[tool.uv.sources] +simpleaudit = { path = "../SimpleAudit", editable = true } + [project.scripts] simpleaudit-studio = "simpleaudit_studio.cli:main" spin = "simpleaudit_studio.cli:main" @@ -56,6 +75,7 @@ packages = [ "config", "accounts", "audits", + "chat", "infra", "scenarios", "model_registry", diff --git a/pytest.ini b/pytest.ini new file mode 100644 index 00000000..f9e972c7 --- /dev/null +++ b/pytest.ini @@ -0,0 +1,20 @@ +[pytest] +; Parsed before conftest.py, so these reach Django settings on import. +; Mirrors the env block in .github/workflows/ci.yml. +DJANGO_SETTINGS_MODULE = config.settings +; e2e/test_ui.py is a standalone Playwright script (not a pytest test) that +; sys.exit()s at import, so it must not be collected. +norecursedirs = e2e node_modules .venv static staticfiles +env = + SIMPLEAUDIT_LOCAL_SQLITE=1 + DJANGO_SECRET_KEY=test-secret-key-not-change-me + DJANGO_DEBUG=false + DJANGO_ALLOWED_HOSTS=* +; Django @tag("...") values are mirrored onto pytest items by conftest.py, so +; -m selection works the same as manage.py test --exclude-tag. +markers = + slow: heavy integration tests (engine, full API lifecycle, monitors, judges, experiments) + embedded_hatchet: starts a real embedded Hatchet worker; run serially, not in parallel +addopts = -ra --strict-markers +filterwarnings = + ignore::DeprecationWarning diff --git a/simpleaudit_studio/cli.py b/simpleaudit_studio/cli.py index 72fdf926..be41b887 100644 --- a/simpleaudit_studio/cli.py +++ b/simpleaudit_studio/cli.py @@ -12,11 +12,17 @@ uvx simpleaudit-studio # full stack; models point at OpenAI (add your key in the UI) uvx simpleaudit-studio --mock # use the built-in mock model server (zero-setup demo) uvx simpleaudit-studio --port 9000 # custom port + +When a needed port is held by another Studio instance or one of its +derivatives (Open WebUI, the Hatchet sidecar), the CLI offers to stop it and +spin cleanly. --no-force-kill declines that offer (and exits with a free-port +suggestion); --yes skips the confirmation. """ from __future__ import annotations import os +import secrets import signal import threading import time @@ -38,18 +44,40 @@ def main() -> None: "--mock", action="store_true", help="Use the built-in mock model server (zero-setup demo; results are simulated)", ) + parser.add_argument( + "--disable-chat", "--no-chat", dest="disable_chat", action="store_true", + help="Do not run the bundled chat (Open WebUI); /chat/ stays unavailable", + ) parser.add_argument( "--no-browser", action="store_true", help="Do not auto-open the web UI in the default browser", ) + parser.add_argument( + "--no-force-kill", action="store_true", + help="Never stop another Studio instance or its derivatives to free a " + "port; exit with a free-port suggestion instead", + ) + parser.add_argument( + "--yes", action="store_true", + help="Answer yes to the force-kill confirmation without asking", + ) args = parser.parse_args() # Set local mode BEFORE Django reads settings os.environ["SIMPLEAUDIT_MINIMAL"] = "1" + # Chat is part of the bundle; --disable-chat (or SIMPLEAUDIT_CHAT=disabled) + # opts out. + if args.disable_chat: + os.environ["SIMPLEAUDIT_CHAT"] = "off" + else: + os.environ.setdefault("SIMPLEAUDIT_CHAT", "embedded") os.environ.setdefault("DJANGO_SETTINGS_MODULE", "config.settings") os.environ.setdefault("DJANGO_SECRET_KEY", "local-insecure-key-change-for-shared-use") os.environ.setdefault("DJANGO_ALLOWED_HOSTS", "localhost,127.0.0.1") - os.environ.setdefault("DJANGO_DEBUG", "true") + # Local dev runs with DEBUG off by default so you experience the real + # production behavior (host checks, static serving, no debug error pages). + # Set DJANGO_DEBUG=true to get the debug tooling back. + os.environ.setdefault("DJANGO_DEBUG", "false") import django @@ -83,7 +111,7 @@ def main() -> None: print("✅ Demo data ready.\n") # --- Step 3: Pre-check port availability --- - _check_port_available(args.port) + _ensure_ports_available(args) # --- Step 4: Start embedded Hatchet --- from infra.minimal_config import start_embedded_hatchet, stop_embedded_hatchet @@ -121,10 +149,37 @@ def main() -> None: # Give the web server a moment to bind time.sleep(1) + # --- Chat: Open WebUI + its forward-auth proxy --- + chat_process = None + from chat import config as chat_config + + if chat_config.ENABLED: + from chat import proxy as chat_proxy + + print("💬 Starting chat (Open WebUI)...") + if chat_proxy.is_first_run(): + print(" First start downloads it (~1 GB via uvx) and can take a few minutes.") + print(" Studio is usable right away; /chat/ works once the download finishes.") + print(f" Its data: {chat_proxy.home_dir()}") + print(f" Its log: {chat_proxy.log_path()}") + try: + chat_process = start_chat(chat_proxy, port) + except (OSError, RuntimeError) as exc: + # No open-webui to run, a taken port, a failed spawn: chat is one + # part of the stack, so the rest still comes up without it. + print(f"⚠️ Chat could not start ({exc}); continuing without it.") + print(" Skip it with --disable-chat.\n") + else: + print() + username = os.environ.get("BOOTSTRAP_USERNAME", "studio") password = os.environ.get("BOOTSTRAP_PASSWORD", "admin123") - auto_login_url = f"http://localhost:{port}/auto-login/" + # One-time sign-in token: the browser opens the URL below, and the first + # request that presents it is the only one that works (see auto_login_view). + auto_login_token = secrets.token_urlsafe(32) + os.environ["SIMPLEAUDIT_AUTO_LOGIN_TOKEN"] = auto_login_token + auto_login_url = f"http://localhost:{port}/auto-login/?token={auto_login_token}" print("┌─────────────────────────────────────────────────────────┐") print("│ │") print("│ 🚀 SimpleAudit Studio is running! │") @@ -132,6 +187,8 @@ def main() -> None: print(f"│ Web UI: http://localhost:{port} │") print(f"│ Login: {username} / {password:<20s}│") print(f"│ API Docs: http://localhost:{port}/api/schema/ │") + if chat_process is not None: + print(f"│ Chat: http://localhost:{port}/chat/ │") print("│ │") if args.mock: print("│ Models: Built-in mock (simulated results) │") @@ -141,12 +198,16 @@ def main() -> None: print("│ Press Ctrl+C to stop. │") print("└─────────────────────────────────────────────────────────┘") print() + # Single-use sign-in link: opens the browser signed in; if no browser + # opens, the user can paste this URL manually (it works exactly once). + print(f"🔑 One-time sign-in link: {auto_login_url}") + print() # Open the default browser signed-in (non-fatal; skip with --no-browser). if not args.no_browser: threading.Thread( target=_open_browser_when_ready, - args=(auto_login_url,), + args=(auto_login_url, port), daemon=True, ).start() @@ -169,38 +230,110 @@ def _sigterm_handler(signum, frame): try: stop_embedded_hatchet() finally: + if chat_process is not None: + from chat.proxy import stop_open_webui + + stop_open_webui() if mock_server is not None: mock_server.shutdown() -def _check_port_available(port: int) -> None: - """Exit early with a clear message if the web port is already in use.""" - import socket +def start_chat(chat_proxy, studio_port: int): + """Start Open WebUI and its proxy, and report readiness in the background. - with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as s: - s.setsockopt(socket.SOL_SOCKET, socket.SO_REUSEADDR, 1) - try: - s.bind(("0.0.0.0", port)) - except OSError: - print(f"\n✗ Port {port} is already in use.") - print(" Is another SimpleAudit Studio instance running?") - print(f" Try a different port: spin --port {port + 1}\n") - raise SystemExit(1) + Open WebUI takes minutes to be ready on a first run (it is fetched, then it + migrates its database), so the wait happens in a thread: Studio and the + worker come up meanwhile, and one line says when /chat/ is live. + """ + process = chat_proxy.start_open_webui(studio_port) + chat_proxy.serve(studio_port) + + def report(): + # flush: this lands minutes later, and stdout is block-buffered when the + # CLI's output is a file or a pipe rather than a terminal. + if chat_proxy.wait_until_ready(process): + print(f"\n✅ Chat is ready — http://localhost:{studio_port}/chat/", flush=True) + print(f" {_sync_chat_models()}\n", flush=True) + elif process.poll() is not None: + print(f"\n⚠️ Chat stopped (exit {process.returncode}). Studio is unaffected.") + print(f" What happened: {chat_proxy.log_path()}\n", flush=True) + else: + print("\n⚠️ Chat is still not answering. Studio is unaffected.") + print(f" What it is doing: {chat_proxy.log_path()}\n", flush=True) + + threading.Thread(target=report, daemon=True).start() + return process + + +def _sync_chat_models() -> str: + """Give the fresh chat Studio's model connections, and say how it went. + + Signals keep it in step afterwards (chat/signals.py); this is the first one, + for a chat that has just started or was off while connections changed. + """ + from chat.api import ChatAPIError + from chat.sync import push_now + + try: + result = push_now() + except ChatAPIError as exc: + return f"Models not synced to chat: {exc}" + kept = f", kept {result['kept']} added in chat" if result["kept"] else "" + return f"Synced {result['pushed']} model connection(s) to chat{kept}." + + +def _ensure_ports_available(args) -> None: + """Make sure every port the run needs is free, or exit with advice. + + The web port always; with chat, Open WebUI's upstream and the proxy's too. + A port held by another Studio instance or one of its derivatives can be + stopped (offered, or automatic with --yes); anything else — or a declined + offer — ends the run with a free port and the exact flag or environment + variable that moves the conflicting one. + """ + from urllib.parse import urlsplit + + from simpleaudit_studio import ports + + force_kill = not args.no_force_kill + yes = args.yes + + ports.resolve_port_conflict( + args.port, "the web server", "spin --port ", + force_kill=force_kill, yes=yes, + ) + + from chat import config as chat_config + + if not chat_config.ENABLED: + return + upstream_port = urlsplit(chat_config.UPSTREAM).port or 8080 + ports.resolve_port_conflict( + upstream_port, "chat's Open WebUI", + "SIMPLEAUDIT_CHAT_UPSTREAM=http://127.0.0.1:", + force_kill=force_kill, yes=yes, + ) + ports.resolve_port_conflict( + chat_config.PROXY_PORT, "chat's proxy", + "SIMPLEAUDIT_CHAT_PROXY_PORT=", + force_kill=force_kill, yes=yes, + ) -def _open_browser_when_ready(url: str, timeout: float = 30.0) -> None: +def _open_browser_when_ready(url: str, port: int, timeout: float = 30.0) -> None: """Wait until the web server answers, then open `url` in the default browser. - The URL is /auto-login/, which signs the visitor in and redirects to the - dashboard. Runs in a daemon thread; any failure just prints a hint. + The URL is the one-time /auto-login/?token=... link, which signs the + visitor in and redirects to the dashboard. Readiness is polled on /healthz + (unauthenticated) so the single-use token is not consumed by the probe. + Runs in a daemon thread; any failure just prints a hint. """ import requests deadline = time.monotonic() + timeout while time.monotonic() < deadline: try: - # allow_redirects=False: we only need proof the endpoint exists. - if requests.get(url, timeout=2).status_code in (200, 302): + if requests.get(f"http://localhost:{port}/healthz", timeout=2).status_code == 200: break except requests.RequestException: time.sleep(0.5) diff --git a/simpleaudit_studio/ports.py b/simpleaudit_studio/ports.py new file mode 100644 index 00000000..96ea7727 --- /dev/null +++ b/simpleaudit_studio/ports.py @@ -0,0 +1,245 @@ +"""Who is sitting on a port, and what to do about it. + +The CLI needs a few ports (the web port, and — with chat — Open WebUI's +upstream and the proxy's). When one is taken, the useful answer depends on +*who* is taking it: + +- another SimpleAudit Studio instance, or something it spawned (Open WebUI, + the Hatchet sidecar) — ours to manage: offer to stop it and spin cleanly; +- anything else — not ours to touch: suggest a free port and the exact flag + or environment variable that changes the conflicting one. + +Port ownership and the process tree are read through ``psutil`` rather than +``lsof``/``ps``: the same ``proc_pidinfo`` restriction that makes ``lsof`` +blind to other processes' sockets in some sandboxes also makes the *global* +``psutil.net_connections()`` call raise, but the per-process +``Process(pid).connections()`` call works everywhere, so we scan PIDs one at +a time and skip the ones we are not allowed to inspect. +""" +from __future__ import annotations + +import os +import signal +import socket +import time +from dataclasses import dataclass + +import psutil + +# Command-line markers. A Studio run's listener is the python process running +# the entry script (``.../bin/spin``, ``.../bin/simpleaudit-studio``) or +# ``manage.py runserver``; its children carry the derivative's name. The +# package name is matched with its underscore spelling on purpose: the +# hyphenated project name also appears in checkout paths, where it must not +# count. +_STUDIO_MARKERS = ("simpleaudit_studio", "manage.py runserver") +_STUDIO_SCRIPTS = ("/spin", "/simpleaudit-studio") +_OPEN_WEBUI_MARKER = "open-webui" +_HATCHET_MARKER = "hatchet-embedded-sidecar" + +#: Occupants the CLI may stop on the user's behalf. +OURS = ("studio", "open_webui", "hatchet_sidecar") + + +@dataclass(frozen=True) +class PortOwner: + """The process listening on a port, and what it looks like.""" + pid: int + kind: str # "studio" | "open_webui" | "hatchet_sidecar" | "unknown" + command: str # full command line, for display + + +def _listening_pid(port: int) -> int | None: + """The PID listening on ``port``, or None when the port is free. + + Scans PIDs one at a time. The global ``psutil.net_connections()`` call + walks every process and dies on the first one macOS refuses to inspect + (``proc_pidinfo`` → EPERM on protected system PIDs); the per-process call + does not, so we skip the PIDs we cannot read. + """ + for pid in psutil.pids(): + try: + for conn in psutil.Process(pid).connections(kind="inet"): + if ( + conn.laddr is not None + and conn.laddr[1] == port + and conn.status == psutil.CONN_LISTEN + ): + return pid + except (psutil.AccessDenied, psutil.NoSuchProcess, psutil.ZombieProcess): + continue + return None + + +def find_port_owner(port: int) -> PortOwner | None: + """The process listening on ``port``, or None when the port is free.""" + pid = _listening_pid(port) + if pid is None: + return None + return classify_port_owner(pid) + + +def classify_port_owner(pid: int) -> PortOwner | None: + """What the process on a port is, from its command line. + + None when the PID is gone or its command line cannot be read. + """ + try: + cmdline = psutil.Process(pid).cmdline() + except (psutil.NoSuchProcess, psutil.AccessDenied, psutil.ZombieProcess): + return None + if not cmdline: + return None + command = " ".join(cmdline) + return PortOwner(pid=pid, kind=_classify_command(command), command=command) + + +def _classify_command(command: str) -> str: + lowered = command.lower() + if any(marker in lowered for marker in _STUDIO_MARKERS): + return "studio" + if any(token.endswith(script) for script in _STUDIO_SCRIPTS + for token in command.split()): + return "studio" + if _OPEN_WEBUI_MARKER in lowered: + return "open_webui" + if _HATCHET_MARKER in lowered: + return "hatchet_sidecar" + return "unknown" + + +def _children_of(pid: int) -> list[int]: + try: + return [child.pid for child in psutil.Process(pid).children()] + except (psutil.NoSuchProcess, psutil.AccessDenied, psutil.ZombieProcess): + return [] + + +def _kill_tree(pid: int, sig: signal.Signals) -> None: + """Signal ``pid`` and everything it spawned, deepest first. + + A Studio run's children (Open WebUI, the Hatchet sidecar) are not in its + process group, so signalling the group would miss them; walk the tree + instead. + """ + tree: list[int] = [] + pending = [pid] + while pending: + current = pending.pop() + tree.append(current) + pending.extend(_children_of(current)) + for member in reversed(tree): + try: + os.kill(member, sig) + except (ProcessLookupError, PermissionError): + pass + + +def kill_port_owner(owner: PortOwner, grace: float = 10.0) -> bool: + """Stop the occupant and its descendants. True when they are all gone.""" + _kill_tree(owner.pid, signal.SIGTERM) + deadline = time.monotonic() + grace + while time.monotonic() < deadline: + if not _pid_alive(owner.pid): + break + time.sleep(0.2) + if _pid_alive(owner.pid): + _kill_tree(owner.pid, signal.SIGKILL) + deadline = time.monotonic() + 5.0 + while time.monotonic() < deadline and _pid_alive(owner.pid): + time.sleep(0.2) + return not _pid_alive(owner.pid) + + +def _pid_alive(pid: int) -> bool: + """True when the process is still running. + + A zombie counts as gone: it is a dead process awaiting reaping and holds + no ports, so for "is the occupant still there" it is not. (``os.kill(pid, + 0)`` and even ``psutil.is_running()`` both report a zombie as alive.) + """ + try: + return psutil.Process(pid).status() != psutil.STATUS_ZOMBIE + except psutil.NoSuchProcess: + return False + except psutil.AccessDenied: + return True + + +def _port_free(port: int) -> bool: + with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as s: + s.setsockopt(socket.SOL_SOCKET, socket.SO_REUSEADDR, 1) + try: + s.bind(("0.0.0.0", port)) + except OSError: + return False + return True + + +def find_free_port(start: int, limit: int = 100) -> int | None: + """The first free port at or above ``start``, or None after ``limit``.""" + for port in range(start, start + limit): + if _port_free(port): + return port + return None + + +def resolve_port_conflict( + port: int, + what: str, + hint: str, + *, + force_kill: bool, + yes: bool, +) -> int: + """Make sure ``port`` is free, or exit with the best advice. + + ``what`` is what the port is for ("the web server", "chat's Open WebUI"), + ``hint`` is the exact flag or environment variable that moves it. When the + occupant is another Studio instance or one of its derivatives, stopping it + is offered (and done outright with ``--yes``); with ``--no-force-kill`` — + or when the occupant is not ours — the run exits and points at a free port + instead. + """ + owner = find_port_owner(port) + if owner is None: + return port + if owner.kind in OURS: + if not force_kill: + _exit_with_free_port(port, what, hint, owner) + label = { + "studio": "another SimpleAudit Studio instance", + "open_webui": "an Open WebUI spawned by SimpleAudit Studio", + "hatchet_sidecar": "a Hatchet sidecar spawned by SimpleAudit Studio", + }[owner.kind] + print(f"\n⚠ Port {port} is held by {label} (PID {owner.pid}):") + print(f" {owner.command}") + if yes: + confirmed = True + else: + try: + answer = input(" Stop it and spin cleanly? [Y/n] ").strip().lower() + except EOFError: + answer = "y" + confirmed = answer in ("", "y", "yes") + if not confirmed: + _exit_with_free_port(port, what, hint, owner) + print(f" Stopping PID {owner.pid} and its children...") + if kill_port_owner(owner): + print(f" Port {port} is free. Continuing.") + return port + print(f"\n✗ Could not stop PID {owner.pid}.") + print(f" Move {what}: {hint}\n") + raise SystemExit(1) + _exit_with_free_port(port, what, hint, owner) + + +def _exit_with_free_port(port: int, what: str, hint: str, owner: PortOwner | None) -> None: + free = find_free_port(port + 1) + print(f"\n✗ Port {port} is in use by {what}.") + if owner is not None: + print(f" PID {owner.pid}: {owner.command}") + if free is not None: + print(f" A free port: {free} — move {what} with: {hint}") + print() + raise SystemExit(1) diff --git a/simpleaudit_studio/tests/__init__.py b/simpleaudit_studio/tests/__init__.py new file mode 100644 index 00000000..e69de29b diff --git a/simpleaudit_studio/tests/test_ports.py b/simpleaudit_studio/tests/test_ports.py new file mode 100644 index 00000000..96d9b154 --- /dev/null +++ b/simpleaudit_studio/tests/test_ports.py @@ -0,0 +1,243 @@ +"""Port-occupant detection and the force-kill / free-port decision. + +Run: + SIMPLEAUDIT_LOCAL_SQLITE=1 uv run manage.py test simpleaudit_studio.tests.test_ports + +Every listener is a *child process*: the test runner's own command line +contains ``simpleaudit_studio`` (``manage.py test simpleaudit_studio...``), +so an in-process socket would be misclassified as a Studio run. +""" +import socket +import subprocess +import sys +import time + +from django.test import SimpleTestCase + +from simpleaudit_studio import ports + + +def _free_port() -> int: + with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as s: + s.bind(("127.0.0.1", 0)) + return s.getsockname()[1] + + +_LISTENER_SCRIPT = ( + "import socket, time\n" + "s = socket.socket()\n" + "s.setsockopt(socket.SOL_SOCKET, socket.SO_REUSEADDR, 1)\n" + "s.bind(('127.0.0.1', {port}))\n" + "s.listen(1)\n" + "time.sleep(60)\n" +) + + +def _listener_proc(port: int, marker: str = "") -> subprocess.Popen: + """A child process holding ``port``; ``marker`` is put in its command + line so ``_classify_command`` sees it (the way a real run's does).""" + script = _LISTENER_SCRIPT.format(port=port) + if marker: + script = f"# {marker}\n" + script + return subprocess.Popen( + [sys.executable, "-c", script], + stdout=subprocess.DEVNULL, stderr=subprocess.DEVNULL, + ) + + +def _wait_port_owned(port: int, timeout: float = 15.0): + deadline = time.monotonic() + timeout + while time.monotonic() < deadline: + owner = ports.find_port_owner(port) + if owner is not None: + return owner + time.sleep(0.1) + raise AssertionError(f"port {port} never became occupied") + + +def _wait_port_free(port: int, timeout: float = 15.0) -> None: + deadline = time.monotonic() + timeout + while time.monotonic() < deadline: + if ports.find_port_owner(port) is None: + return + time.sleep(0.1) + raise AssertionError(f"port {port} still occupied after {timeout}s") + + +class FindPortOwnerTests(SimpleTestCase): + def test_a_free_port_has_no_owner(self): + self.assertIsNone(ports.find_port_owner(_free_port())) + + def test_a_plain_socket_is_unknown(self): + port = _free_port() + proc = _listener_proc(port) + try: + owner = _wait_port_owned(port) + self.assertEqual(owner.kind, "unknown") + self.assertEqual(owner.pid, proc.pid) + finally: + proc.kill() + proc.wait() + + def test_a_studio_process_is_classified(self): + port = _free_port() + proc = _listener_proc(port, marker="simpleaudit_studio") + try: + owner = _wait_port_owned(port) + self.assertEqual(owner.kind, "studio") + self.assertEqual(owner.pid, proc.pid) + finally: + proc.kill() + proc.wait() + + def test_an_open_webui_process_is_classified(self): + port = _free_port() + proc = _listener_proc(port, marker="open-webui") + try: + owner = _wait_port_owned(port) + self.assertEqual(owner.kind, "open_webui") + finally: + proc.kill() + proc.wait() + + def test_a_hatchet_sidecar_is_classified(self): + port = _free_port() + proc = _listener_proc(port, marker="hatchet-embedded-sidecar") + try: + owner = _wait_port_owned(port) + self.assertEqual(owner.kind, "hatchet_sidecar") + finally: + proc.kill() + proc.wait() + + +class KillPortOwnerTests(SimpleTestCase): + def test_kills_the_process_and_its_children(self): + port = _free_port() + # A parent that spawns a child which is not in its process group — + # the shape of a Studio run (its Open WebUI / sidecar children). + proc = subprocess.Popen( + [sys.executable, "-c", + ("import socket, subprocess, sys, time\n" + "s = socket.socket()\n" + "s.setsockopt(socket.SOL_SOCKET, socket.SO_REUSEADDR, 1)\n" + f"s.bind(('127.0.0.1', {port})); s.listen(1)\n" + "c = subprocess.Popen([sys.executable, '-c', " + "'import time; time.sleep(60)'], start_new_session=True)\n" + "time.sleep(60)")], + stdout=subprocess.DEVNULL, stderr=subprocess.DEVNULL, + ) + try: + owner = _wait_port_owned(port) + children = ports._children_of(proc.pid) + self.assertEqual(len(children), 1) + self.assertTrue(ports.kill_port_owner(owner)) + _wait_port_free(port) + self.assertFalse(ports._pid_alive(children[0])) + finally: + proc.kill() + proc.wait() + + +class ResolvePortConflictTests(SimpleTestCase): + def _studio_occupant(self, port: int) -> subprocess.Popen: + return _listener_proc(port, marker="simpleaudit_studio") + + def test_a_free_port_passes_through(self): + port = _free_port() + self.assertEqual( + ports.resolve_port_conflict(port, "the web server", "spin --port N", + force_kill=True, yes=True), + port, + ) + + def test_studio_occupant_is_killed_with_yes(self): + port = _free_port() + proc = self._studio_occupant(port) + try: + _wait_port_owned(port) + result = ports.resolve_port_conflict( + port, "the web server", "spin --port N", + force_kill=True, yes=True, + ) + self.assertEqual(result, port) + _wait_port_free(port) + finally: + proc.kill() + proc.wait() + + def test_no_force_kill_exits_with_a_free_port(self): + port = _free_port() + proc = self._studio_occupant(port) + try: + _wait_port_owned(port) + with self.assertRaises(SystemExit) as ctx: + ports.resolve_port_conflict( + port, "the web server", "spin --port N", + force_kill=False, yes=True, + ) + self.assertEqual(ctx.exception.code, 1) + self.assertTrue(ports._pid_alive(proc.pid)) # left alone + finally: + proc.kill() + proc.wait() + + def test_a_declined_offer_exits(self): + port = _free_port() + proc = self._studio_occupant(port) + try: + _wait_port_owned(port) + with self._patched_input("n"), self.assertRaises(SystemExit) as ctx: + ports.resolve_port_conflict( + port, "the web server", "spin --port N", + force_kill=True, yes=False, + ) + self.assertEqual(ctx.exception.code, 1) + self.assertTrue(ports._pid_alive(proc.pid)) + finally: + proc.kill() + proc.wait() + + def test_an_accepted_offer_kills(self): + port = _free_port() + proc = self._studio_occupant(port) + try: + _wait_port_owned(port) + with self._patched_input("y"): + result = ports.resolve_port_conflict( + port, "the web server", "spin --port N", + force_kill=True, yes=False, + ) + self.assertEqual(result, port) + _wait_port_free(port) + finally: + proc.kill() + proc.wait() + + def test_an_unknown_occupant_is_never_killed(self): + port = _free_port() + proc = _listener_proc(port) # no marker — not ours + try: + _wait_port_owned(port) + with self.assertRaises(SystemExit) as ctx: + ports.resolve_port_conflict( + port, "the web server", "spin --port N", + force_kill=True, yes=True, + ) + self.assertEqual(ctx.exception.code, 1) + self.assertTrue(ports._pid_alive(proc.pid)) # not ours to touch + finally: + proc.kill() + proc.wait() + + @staticmethod + def _patched_input(answer: str): + from unittest.mock import patch + return patch("builtins.input", return_value=answer) + + +class FindFreePortTests(SimpleTestCase): + def test_it_returns_a_free_port(self): + port = ports.find_free_port(_free_port()) + self.assertIsNotNone(port) + self.assertTrue(ports._port_free(port)) diff --git a/templates/base.html b/templates/base.html index 465847a4..95ae589a 100644 --- a/templates/base.html +++ b/templates/base.html @@ -72,7 +72,7 @@ {% for item in nav_items %}{% if item.active %}{{ item.label }}{% endif %}{% empty %}SimpleAudit{% endfor %} -
+
{% if messages %}
diff --git a/templates/connections.html b/templates/connections.html index c3f019b6..cc6e3ecf 100644 --- a/templates/connections.html +++ b/templates/connections.html @@ -64,6 +64,10 @@

{{ conn.name }}

{% if conn.can_edit %} + {% if otlp_enabled %} + + {% endif %} {% if conn.in_use %} @@ -122,6 +126,15 @@

{{ conn.name }}

title="Runs and monitors that use this model as target, auditor or judge">{% if m.usage %}{{ m.usage }} run{{ m.usage|pluralize }}{% else %}unused{% endif %} {% if conn.can_edit %} + {% if chat_enabled %} + + + + + + {% endif %} @@ -302,6 +315,33 @@

Add models to +
+
+

OTLP credential for

+ +
+

Issue a credential so this target can push OTLP traces to Studio. The secret is shown once — copy it now.

+
+ + + +
+ +
+ + {{ provider_presets|json_script:"provider-presets" }} {{ conn_data|json_script:"conn-data" }}