mirror of
https://github.com/turnstonelabs/turnstone.git
synced 2026-08-12 23:12:23 -06:00
f4fd7e1f67
* fix(security): classify outbound addresses by what they reach (GHSA-wm4f-79pw-pfr9)
Five guards screened outbound URLs and each hand-rolled its own address
normalization and policy tests, so each had a different hole. An IPv6
transition address carries an IPv4 destination in its low bits and
`ipaddress` classifies the wrapper, not the destination: 64:ff9b::a9fe:a9fe
reports is_global because 64:ff9b::/96 is global unicast, while a NAT64
gateway routes it to the cloud metadata endpoint. CGNAT (100.64.0.0/10) is
neither is_private nor is_global, so a denylist built on is_private missed
it with no gateway involved at all.
Add turnstone/core/ip_classify.py as the single classifier. One function
returns exactly one policy lane — PUBLIC, PRIVATE (operator-approvable) or
NEVER — and every guard branches on the lane rather than re-deriving it.
Two overlapping booleans would make a verdict depend on which one a caller
tested first; several addresses are simultaneously globally routable and
metadata-reaching.
- Decode transition addresses per RFC 6052 §2.2 (NAT64 well-known and
local-use prefixes, 6to4, Teredo, IPv4-mapped, IPv4-compatible) and judge
them by the IPv4 they reach. The local-use prefix does not say which
layout its gateway uses, so every length it can carry is decoded and the
worst result classified.
- Share hostname resolution too. The five copies had already drifted on
which failures they caught, and getaddrinfo raises UnicodeError — not an
OSError — from the IDNA encoder.
- Resolution failure is a refusal, not a pass: the fetch resolves again, so
an authority answering the guard with SERVFAIL and the fetch with an
internal address would otherwise switch the guard off for that hop.
- Screen every redirect hop in every mode. allow_private_origin widens which
lanes are acceptable rather than turning screening off, and the permission
is revoked after any hop that is not wholly private.
- Cleartext http is allowed only for a hostname that RESOLVES to loopback.
*.localhost is ordinary DNS, and trusting the name put an OIDC token
exchange on the wire in the clear.
- Screen doctor and console-probe URLs through the classifier. Both used a
host.startswith("169.254.") string test that never resolved, so any DNS
name pointing at the metadata service passed and its body was returned to
the model.
- Add known vendor metadata prefixes the stdlib does not flag, and place
deprecated IPv6 site-local outside the public lane.
The operator's private-network opt-in still admits the whole home lab,
including IPv6 loopback, CGNAT and split-horizon hosts. Metadata,
link-local, multicast, unspecified and reserved addresses stay refused
regardless of the opt-in, including as a redirect target from an approved
private origin — the settings help and docs now say so.
Reported by @tonghuaroot.
* fix(security): close Azure/Oracle metadata gap and restore dual-stack origins
Review follow-ups on the address-classification rework.
Azure's host-agent endpoint (168.63.129.16) and Oracle Cloud's metadata
endpoint (192.0.0.192) sit in ordinary unicast space, so the stdlib reported
them as globally routable and both classified PUBLIC — reachable with no
opt-in at all, a worse position than the RFC 1918 host beside them, and
directly contradicting the "metadata stays refused even with the opt-in"
guarantee the settings help and docs now advertise. Both join the shared
vendor list.
Revoking the private-hop permission on the ORIGIN hop broke the case
`_screen_tool_url` deliberately admits: a dual-stack or split-horizon
home-lab host answering with both a LAN and a public record was approved,
then refused on its own `302 /login` — one hop was all it ever got. Track
the approved HOST instead, so redirects that stay on it remain covered while
a redirect to any other private host is still refused once the chain is no
longer wholly private.
Also:
- Try several registry candidates for the collector-scope probe instead of
abandoning it when the first is unresolvable, which also stopped a healthy
registry from logging as malformed.
- Bound the probe's name resolution with an explicit timeout matching the
2s the httpx connect deadline used to provide; it runs before the console
lifespan yields and getaddrinfo has no timeout of its own.
- Route doctor and the console probe through `web.screen_url` rather than
keeping a third and fourth copy of parse/resolve/classify/fold, which had
already diverged on default port and empty-hostname wording. An empty
hostname no longer reports as a cloud-metadata refusal.
- Give `screen_url` a scheme-aware default port.
- Stop doubling the word "hostname" in the OAuth resolution refusal.
- Correct the `_screen_tool_url` docstring: it described `private_origin` as
requiring every record to be private, which the mixed-record decision
reversed, and `private_block` as a property of a refusal when it reports
the lane on the success path too.
- Make the preview tests' screening stub opt-in rather than autouse — as a
module-wide fixture it also stubbed the tests whose subject IS the screen,
so one of them would have passed even if screening refused everything.
Verified the module now passes with all name resolution blocked.
* fix(security): refuse mixed-record private origins instead of exempting them
The previous commit let an approved private origin redirect to itself by
exempting its hostname from the chain-wide revocation. That exemption was
wrong three ways: it was captured once and never cleared, so a public hop
could steer the fetcher back into the approved host at a path of its
choosing — reopening the private -> public -> private bypass; it was
re-entrant across same-host redirects with fresh DNS each time, so a
self-redirecting host could walk arbitrary internal addresses; and it
matched on bare hostname, so it spanned every port on the approved box.
All three were reproduced against the parent commit, which refuses them.
Delete the exemption rather than repair it. The case it existed for — a
dual-stack host answering with both a LAN and a public record — is now
refused where it is actually decidable, in `_screen_tool_url`, with the
remedy in the message: point the tool at the LAN address directly. A
granted chain therefore always starts wholly private, so the fetch guard
needs no notion of an approved host and stays one unconditional rule.
That the accommodation could not be expressed safely in the guard is the
signal: the connection may land on either record, so approving such a host
never described where the fetch would go.
Also from the same review:
- Walk the whole service registry for a collector-scope probe candidate
instead of the first three, and split the outcome into three log lines,
so entries that are merely unreachable stop raising the malformed-registry
alarm and skipping the boot check cluster-wide.
- Stop the candidate walk on a resolver timeout. `asyncio.timeout` bounds
the await, not the work, so continuing left one parked thread per timed-out
candidate on the shared executor.
- Move the metadata-hostname denylist into `ip_classify` and enforce it in
`screen_url`, so doctor and the console probe inherit it instead of each
keeping a copy.
- Drop the scheme-aware default port: a numeric service does not change
which addresses resolution returns, and classification reads only those.
`parsed.port` is still touched so an out-of-range value refuses.
- Correct the vendor-metadata comment, which generalized a claim true of
Azure's and Oracle's addresses to Alibaba's CGNAT one.
- Rename a test class that was still named for the rule it no longer tests.
2928 lines
115 KiB
Python
2928 lines
115 KiB
Python
"""Tests for turnstone.core.oidc — OIDC authentication support."""
|
|
|
|
from __future__ import annotations
|
|
|
|
import asyncio
|
|
import base64
|
|
import dataclasses
|
|
import hashlib
|
|
import types
|
|
import urllib.parse
|
|
from types import SimpleNamespace
|
|
from typing import Any
|
|
from unittest.mock import AsyncMock, MagicMock, patch
|
|
|
|
import httpx
|
|
import jwt as pyjwt
|
|
import pytest
|
|
|
|
from tests.conftest import make_oidc_test_config as _make_config
|
|
from turnstone.core.oidc import (
|
|
OIDCError,
|
|
OIDCKeyNotFoundError,
|
|
_ensure_default_role,
|
|
_sanitize_log_text,
|
|
apply_role_mapping,
|
|
build_authorize_url,
|
|
discover_oidc,
|
|
exchange_code,
|
|
generate_pkce_verifier,
|
|
initialize_oidc_state,
|
|
load_oidc_config,
|
|
provision_oidc_user,
|
|
validate_discovered_endpoint,
|
|
validate_id_token,
|
|
validate_issuer_url,
|
|
)
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# Helpers
|
|
# ---------------------------------------------------------------------------
|
|
|
|
|
|
def _mock_storage(**overrides):
|
|
"""Build a MagicMock with sensible storage defaults."""
|
|
s = MagicMock()
|
|
s.get_oidc_identity.return_value = overrides.get("identity")
|
|
s.get_user.return_value = overrides.get("user")
|
|
s.get_user_by_username.return_value = overrides.get("user_by_username")
|
|
s.get_role.return_value = overrides.get("role")
|
|
s.find_existing_usernames.return_value = overrides.get("existing_usernames", set())
|
|
s.replace_oidc_roles.return_value = overrides.get("replace_oidc_roles", (set(), set()))
|
|
return s
|
|
|
|
|
|
def _mock_async_client(mock_get):
|
|
"""Build a patched httpx.AsyncClient context manager for async tests."""
|
|
|
|
class _AsyncCtx:
|
|
async def __aenter__(self):
|
|
return self
|
|
|
|
async def __aexit__(self, *args):
|
|
pass
|
|
|
|
async def get(self, url):
|
|
return await mock_get(url)
|
|
|
|
return _AsyncCtx()
|
|
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# Config Loading
|
|
# ---------------------------------------------------------------------------
|
|
|
|
|
|
class TestLoadOIDCConfig:
|
|
def test_load_oidc_config_from_env(self, monkeypatch):
|
|
monkeypatch.setenv("TURNSTONE_OIDC_ISSUER", "https://auth.example.com")
|
|
monkeypatch.setenv("TURNSTONE_OIDC_CLIENT_ID", "cid")
|
|
monkeypatch.setenv("TURNSTONE_OIDC_CLIENT_SECRET", "csecret")
|
|
monkeypatch.setenv("TURNSTONE_OIDC_SCOPES", "openid")
|
|
monkeypatch.setenv("TURNSTONE_OIDC_PROVIDER_NAME", "Okta")
|
|
|
|
with patch("turnstone.core.config.load_config", return_value={}):
|
|
cfg = load_oidc_config()
|
|
|
|
assert cfg.enabled is True
|
|
assert cfg.issuer == "https://auth.example.com"
|
|
assert cfg.client_id == "cid"
|
|
assert cfg.client_secret == "csecret"
|
|
assert cfg.scopes == "openid"
|
|
assert cfg.provider_name == "Okta"
|
|
|
|
def test_load_oidc_config_allow_private_network_env(self, monkeypatch):
|
|
monkeypatch.setenv("TURNSTONE_OIDC_ISSUER", "https://auth.internal.example")
|
|
monkeypatch.setenv("TURNSTONE_OIDC_CLIENT_ID", "cid")
|
|
monkeypatch.setenv("TURNSTONE_OIDC_CLIENT_SECRET", "csecret")
|
|
monkeypatch.setenv("TURNSTONE_OIDC_ALLOW_PRIVATE_NETWORK", "true")
|
|
|
|
with patch("turnstone.core.config.load_config", return_value={}):
|
|
cfg = load_oidc_config()
|
|
|
|
assert cfg.allow_private_network is True
|
|
|
|
def test_load_oidc_config_allow_private_network_toml(self, monkeypatch):
|
|
monkeypatch.setenv("TURNSTONE_OIDC_ISSUER", "https://auth.internal.example")
|
|
monkeypatch.setenv("TURNSTONE_OIDC_CLIENT_ID", "cid")
|
|
monkeypatch.setenv("TURNSTONE_OIDC_CLIENT_SECRET", "csecret")
|
|
monkeypatch.delenv("TURNSTONE_OIDC_ALLOW_PRIVATE_NETWORK", raising=False)
|
|
|
|
with patch(
|
|
"turnstone.core.config.load_config",
|
|
return_value={"allow_private_network": True},
|
|
):
|
|
cfg = load_oidc_config()
|
|
|
|
assert cfg.allow_private_network is True
|
|
|
|
def test_load_oidc_config_allow_private_network_default_off(self, monkeypatch):
|
|
monkeypatch.setenv("TURNSTONE_OIDC_ISSUER", "https://auth.example.com")
|
|
monkeypatch.setenv("TURNSTONE_OIDC_CLIENT_ID", "cid")
|
|
monkeypatch.setenv("TURNSTONE_OIDC_CLIENT_SECRET", "csecret")
|
|
monkeypatch.delenv("TURNSTONE_OIDC_ALLOW_PRIVATE_NETWORK", raising=False)
|
|
|
|
with patch("turnstone.core.config.load_config", return_value={}):
|
|
cfg = load_oidc_config()
|
|
|
|
assert cfg.allow_private_network is False
|
|
|
|
def test_load_oidc_config_capture_user_credential_env(self, monkeypatch):
|
|
monkeypatch.setenv("TURNSTONE_OIDC_ISSUER", "https://auth.example.com")
|
|
monkeypatch.setenv("TURNSTONE_OIDC_CLIENT_ID", "cid")
|
|
monkeypatch.setenv("TURNSTONE_OIDC_CLIENT_SECRET", "csecret")
|
|
monkeypatch.setenv("TURNSTONE_OIDC_CAPTURE_USER_CREDENTIAL", "true")
|
|
|
|
with patch("turnstone.core.config.load_config", return_value={}):
|
|
cfg = load_oidc_config()
|
|
|
|
assert cfg.capture_user_credential is True
|
|
|
|
def test_load_oidc_config_capture_user_credential_toml(self, monkeypatch):
|
|
monkeypatch.setenv("TURNSTONE_OIDC_ISSUER", "https://auth.example.com")
|
|
monkeypatch.setenv("TURNSTONE_OIDC_CLIENT_ID", "cid")
|
|
monkeypatch.setenv("TURNSTONE_OIDC_CLIENT_SECRET", "csecret")
|
|
monkeypatch.delenv("TURNSTONE_OIDC_CAPTURE_USER_CREDENTIAL", raising=False)
|
|
|
|
with patch(
|
|
"turnstone.core.config.load_config",
|
|
return_value={"capture_user_credential": True},
|
|
):
|
|
cfg = load_oidc_config()
|
|
|
|
assert cfg.capture_user_credential is True
|
|
|
|
def test_load_oidc_config_capture_default_off(self, monkeypatch):
|
|
monkeypatch.setenv("TURNSTONE_OIDC_ISSUER", "https://auth.example.com")
|
|
monkeypatch.setenv("TURNSTONE_OIDC_CLIENT_ID", "cid")
|
|
monkeypatch.setenv("TURNSTONE_OIDC_CLIENT_SECRET", "csecret")
|
|
monkeypatch.delenv("TURNSTONE_OIDC_CAPTURE_USER_CREDENTIAL", raising=False)
|
|
|
|
with patch("turnstone.core.config.load_config", return_value={}):
|
|
cfg = load_oidc_config()
|
|
|
|
assert cfg.capture_user_credential is False
|
|
|
|
def test_capture_appends_offline_access_to_scopes(self, monkeypatch):
|
|
"""Enabling capture requests offline_access without operator scope edits."""
|
|
monkeypatch.setenv("TURNSTONE_OIDC_ISSUER", "https://auth.example.com")
|
|
monkeypatch.setenv("TURNSTONE_OIDC_CLIENT_ID", "cid")
|
|
monkeypatch.setenv("TURNSTONE_OIDC_CLIENT_SECRET", "csecret")
|
|
monkeypatch.delenv("TURNSTONE_OIDC_SCOPES", raising=False)
|
|
|
|
with patch(
|
|
"turnstone.core.config.load_config",
|
|
return_value={"capture_user_credential": True},
|
|
):
|
|
cfg = load_oidc_config()
|
|
|
|
assert cfg.scopes == "openid email profile offline_access"
|
|
|
|
def test_capture_scope_append_is_idempotent(self, monkeypatch):
|
|
"""An operator who already lists offline_access doesn't get it twice."""
|
|
monkeypatch.setenv("TURNSTONE_OIDC_ISSUER", "https://auth.example.com")
|
|
monkeypatch.setenv("TURNSTONE_OIDC_CLIENT_ID", "cid")
|
|
monkeypatch.setenv("TURNSTONE_OIDC_CLIENT_SECRET", "csecret")
|
|
monkeypatch.setenv("TURNSTONE_OIDC_SCOPES", "openid offline_access email")
|
|
|
|
with patch(
|
|
"turnstone.core.config.load_config",
|
|
return_value={"capture_user_credential": True},
|
|
):
|
|
cfg = load_oidc_config()
|
|
|
|
assert cfg.scopes == "openid offline_access email"
|
|
|
|
def test_no_capture_leaves_scopes_untouched(self, monkeypatch):
|
|
monkeypatch.setenv("TURNSTONE_OIDC_ISSUER", "https://auth.example.com")
|
|
monkeypatch.setenv("TURNSTONE_OIDC_CLIENT_ID", "cid")
|
|
monkeypatch.setenv("TURNSTONE_OIDC_CLIENT_SECRET", "csecret")
|
|
monkeypatch.delenv("TURNSTONE_OIDC_SCOPES", raising=False)
|
|
monkeypatch.delenv("TURNSTONE_OIDC_CAPTURE_USER_CREDENTIAL", raising=False)
|
|
|
|
with patch("turnstone.core.config.load_config", return_value={}):
|
|
cfg = load_oidc_config()
|
|
|
|
assert cfg.scopes == "openid email profile"
|
|
|
|
def test_load_oidc_config_disabled_when_missing(self, monkeypatch):
|
|
monkeypatch.delenv("TURNSTONE_OIDC_ISSUER", raising=False)
|
|
monkeypatch.delenv("TURNSTONE_OIDC_CLIENT_ID", raising=False)
|
|
monkeypatch.delenv("TURNSTONE_OIDC_CLIENT_SECRET", raising=False)
|
|
|
|
with patch("turnstone.core.config.load_config", return_value={}):
|
|
cfg = load_oidc_config()
|
|
|
|
assert cfg.enabled is False
|
|
|
|
def test_load_oidc_config_partial_env(self, monkeypatch):
|
|
"""Only issuer set, no client_id -> enabled=False."""
|
|
monkeypatch.setenv("TURNSTONE_OIDC_ISSUER", "https://auth.example.com")
|
|
monkeypatch.delenv("TURNSTONE_OIDC_CLIENT_ID", raising=False)
|
|
monkeypatch.delenv("TURNSTONE_OIDC_CLIENT_SECRET", raising=False)
|
|
|
|
with patch("turnstone.core.config.load_config", return_value={}):
|
|
cfg = load_oidc_config()
|
|
|
|
assert cfg.enabled is False
|
|
assert cfg.issuer == "https://auth.example.com"
|
|
assert cfg.client_id == ""
|
|
|
|
def test_load_oidc_config_role_map_parsing(self, monkeypatch):
|
|
monkeypatch.setenv("TURNSTONE_OIDC_ISSUER", "https://auth.example.com")
|
|
monkeypatch.setenv("TURNSTONE_OIDC_CLIENT_ID", "cid")
|
|
monkeypatch.setenv("TURNSTONE_OIDC_CLIENT_SECRET", "csecret")
|
|
monkeypatch.setenv("TURNSTONE_OIDC_ROLE_CLAIM", "roles")
|
|
monkeypatch.setenv("TURNSTONE_OIDC_ROLE_MAP", "admin:builtin-admin,eng:builtin-operator")
|
|
|
|
with patch("turnstone.core.config.load_config", return_value={}):
|
|
cfg = load_oidc_config()
|
|
|
|
assert cfg.role_claim == "roles"
|
|
assert cfg.role_map == {"admin": "builtin-admin", "eng": "builtin-operator"}
|
|
|
|
def test_load_oidc_config_password_enabled_false(self, monkeypatch):
|
|
monkeypatch.setenv("TURNSTONE_OIDC_ISSUER", "https://auth.example.com")
|
|
monkeypatch.setenv("TURNSTONE_OIDC_CLIENT_ID", "cid")
|
|
monkeypatch.setenv("TURNSTONE_OIDC_CLIENT_SECRET", "csecret")
|
|
monkeypatch.setenv("TURNSTONE_OIDC_PASSWORD_ENABLED", "false")
|
|
|
|
with patch("turnstone.core.config.load_config", return_value={}):
|
|
cfg = load_oidc_config()
|
|
|
|
assert cfg.enabled is True
|
|
assert cfg.password_enabled is False
|
|
|
|
def test_load_oidc_config_password_enabled_true(self, monkeypatch):
|
|
monkeypatch.setenv("TURNSTONE_OIDC_ISSUER", "https://auth.example.com")
|
|
monkeypatch.setenv("TURNSTONE_OIDC_CLIENT_ID", "cid")
|
|
monkeypatch.setenv("TURNSTONE_OIDC_CLIENT_SECRET", "csecret")
|
|
monkeypatch.setenv("TURNSTONE_OIDC_PASSWORD_ENABLED", "true")
|
|
|
|
with patch("turnstone.core.config.load_config", return_value={}):
|
|
cfg = load_oidc_config()
|
|
|
|
assert cfg.password_enabled is True
|
|
|
|
def test_load_oidc_config_role_map_empty_entries(self, monkeypatch):
|
|
"""Role map with empty/whitespace entries should be silently skipped."""
|
|
monkeypatch.setenv("TURNSTONE_OIDC_ISSUER", "https://auth.example.com")
|
|
monkeypatch.setenv("TURNSTONE_OIDC_CLIENT_ID", "cid")
|
|
monkeypatch.setenv("TURNSTONE_OIDC_CLIENT_SECRET", "csecret")
|
|
monkeypatch.setenv("TURNSTONE_OIDC_ROLE_MAP", "admin:builtin-admin, , :, foo:")
|
|
|
|
with patch("turnstone.core.config.load_config", return_value={}):
|
|
cfg = load_oidc_config()
|
|
|
|
assert cfg.role_map == {"admin": "builtin-admin"}
|
|
|
|
def test_load_oidc_config_defaults(self, monkeypatch):
|
|
"""Defaults for scopes and provider_name when not set."""
|
|
monkeypatch.setenv("TURNSTONE_OIDC_ISSUER", "https://auth.example.com")
|
|
monkeypatch.setenv("TURNSTONE_OIDC_CLIENT_ID", "cid")
|
|
monkeypatch.setenv("TURNSTONE_OIDC_CLIENT_SECRET", "csecret")
|
|
monkeypatch.delenv("TURNSTONE_OIDC_SCOPES", raising=False)
|
|
monkeypatch.delenv("TURNSTONE_OIDC_PROVIDER_NAME", raising=False)
|
|
|
|
with patch("turnstone.core.config.load_config", return_value={}):
|
|
cfg = load_oidc_config()
|
|
|
|
assert cfg.scopes == "openid email profile"
|
|
assert cfg.provider_name == "SSO"
|
|
|
|
def test_load_oidc_config_redirect_base_from_env(self, monkeypatch):
|
|
"""TURNSTONE_OIDC_REDIRECT_BASE populates redirect_base."""
|
|
monkeypatch.setenv("TURNSTONE_OIDC_ISSUER", "https://auth.example.com")
|
|
monkeypatch.setenv("TURNSTONE_OIDC_CLIENT_ID", "cid")
|
|
monkeypatch.setenv("TURNSTONE_OIDC_CLIENT_SECRET", "csecret")
|
|
monkeypatch.setenv("TURNSTONE_OIDC_REDIRECT_BASE", "https://app.example.com")
|
|
|
|
with patch("turnstone.core.config.load_config", return_value={}):
|
|
cfg = load_oidc_config()
|
|
|
|
assert cfg.redirect_base == "https://app.example.com"
|
|
|
|
def test_load_oidc_config_redirect_base_strips_trailing_slash(self, monkeypatch):
|
|
"""Trailing slashes are stripped from redirect_base."""
|
|
monkeypatch.setenv("TURNSTONE_OIDC_ISSUER", "https://auth.example.com")
|
|
monkeypatch.setenv("TURNSTONE_OIDC_CLIENT_ID", "cid")
|
|
monkeypatch.setenv("TURNSTONE_OIDC_CLIENT_SECRET", "csecret")
|
|
monkeypatch.setenv("TURNSTONE_OIDC_REDIRECT_BASE", "https://app.example.com/")
|
|
|
|
with patch("turnstone.core.config.load_config", return_value={}):
|
|
cfg = load_oidc_config()
|
|
|
|
assert cfg.redirect_base == "https://app.example.com"
|
|
|
|
def test_load_oidc_config_redirect_base_default_empty(self, monkeypatch):
|
|
"""redirect_base defaults to empty string when not set."""
|
|
monkeypatch.setenv("TURNSTONE_OIDC_ISSUER", "https://auth.example.com")
|
|
monkeypatch.setenv("TURNSTONE_OIDC_CLIENT_ID", "cid")
|
|
monkeypatch.setenv("TURNSTONE_OIDC_CLIENT_SECRET", "csecret")
|
|
monkeypatch.delenv("TURNSTONE_OIDC_REDIRECT_BASE", raising=False)
|
|
|
|
with patch("turnstone.core.config.load_config", return_value={}):
|
|
cfg = load_oidc_config()
|
|
|
|
assert cfg.redirect_base == ""
|
|
|
|
def test_load_oidc_config_redirect_base_rejects_path(self, monkeypatch):
|
|
"""redirect_base with a path component is rejected (falls back to empty)."""
|
|
monkeypatch.setenv("TURNSTONE_OIDC_ISSUER", "https://auth.example.com")
|
|
monkeypatch.setenv("TURNSTONE_OIDC_CLIENT_ID", "cid")
|
|
monkeypatch.setenv("TURNSTONE_OIDC_CLIENT_SECRET", "csecret")
|
|
monkeypatch.setenv("TURNSTONE_OIDC_REDIRECT_BASE", "https://app.example.com/subpath")
|
|
|
|
with patch("turnstone.core.config.load_config", return_value={}):
|
|
cfg = load_oidc_config()
|
|
|
|
assert cfg.redirect_base == ""
|
|
|
|
def test_load_oidc_config_redirect_base_rejects_no_scheme(self, monkeypatch):
|
|
"""redirect_base without a scheme is rejected."""
|
|
monkeypatch.setenv("TURNSTONE_OIDC_ISSUER", "https://auth.example.com")
|
|
monkeypatch.setenv("TURNSTONE_OIDC_CLIENT_ID", "cid")
|
|
monkeypatch.setenv("TURNSTONE_OIDC_CLIENT_SECRET", "csecret")
|
|
monkeypatch.setenv("TURNSTONE_OIDC_REDIRECT_BASE", "app.example.com")
|
|
|
|
with patch("turnstone.core.config.load_config", return_value={}):
|
|
cfg = load_oidc_config()
|
|
|
|
assert cfg.redirect_base == ""
|
|
|
|
def test_load_oidc_config_redirect_base_rejects_userinfo(self, monkeypatch):
|
|
"""redirect_base with userinfo (user:pass@host) is rejected."""
|
|
monkeypatch.setenv("TURNSTONE_OIDC_ISSUER", "https://auth.example.com")
|
|
monkeypatch.setenv("TURNSTONE_OIDC_CLIENT_ID", "cid")
|
|
monkeypatch.setenv("TURNSTONE_OIDC_CLIENT_SECRET", "csecret")
|
|
monkeypatch.setenv("TURNSTONE_OIDC_REDIRECT_BASE", "https://user:pass@app.example.com")
|
|
|
|
with patch("turnstone.core.config.load_config", return_value={}):
|
|
cfg = load_oidc_config()
|
|
|
|
assert cfg.redirect_base == ""
|
|
|
|
def test_load_oidc_config_redirect_base_rejects_invalid_port(self, monkeypatch):
|
|
"""redirect_base with non-numeric port is rejected."""
|
|
monkeypatch.setenv("TURNSTONE_OIDC_ISSUER", "https://auth.example.com")
|
|
monkeypatch.setenv("TURNSTONE_OIDC_CLIENT_ID", "cid")
|
|
monkeypatch.setenv("TURNSTONE_OIDC_CLIENT_SECRET", "csecret")
|
|
monkeypatch.setenv("TURNSTONE_OIDC_REDIRECT_BASE", "https://app.example.com:abc")
|
|
|
|
with patch("turnstone.core.config.load_config", return_value={}):
|
|
cfg = load_oidc_config()
|
|
|
|
assert cfg.redirect_base == ""
|
|
|
|
def test_load_oidc_config_redirect_base_rejects_missing_hostname(self, monkeypatch):
|
|
"""redirect_base without a hostname is rejected."""
|
|
monkeypatch.setenv("TURNSTONE_OIDC_ISSUER", "https://auth.example.com")
|
|
monkeypatch.setenv("TURNSTONE_OIDC_CLIENT_ID", "cid")
|
|
monkeypatch.setenv("TURNSTONE_OIDC_CLIENT_SECRET", "csecret")
|
|
monkeypatch.setenv("TURNSTONE_OIDC_REDIRECT_BASE", "https://")
|
|
|
|
with patch("turnstone.core.config.load_config", return_value={}):
|
|
cfg = load_oidc_config()
|
|
|
|
assert cfg.redirect_base == ""
|
|
|
|
def test_load_oidc_config_redirect_base_allows_http(self, monkeypatch):
|
|
"""http:// redirect_base is allowed (with warning) for local dev."""
|
|
monkeypatch.setenv("TURNSTONE_OIDC_ISSUER", "https://auth.example.com")
|
|
monkeypatch.setenv("TURNSTONE_OIDC_CLIENT_ID", "cid")
|
|
monkeypatch.setenv("TURNSTONE_OIDC_CLIENT_SECRET", "csecret")
|
|
monkeypatch.setenv("TURNSTONE_OIDC_REDIRECT_BASE", "http://localhost:8000")
|
|
|
|
with patch("turnstone.core.config.load_config", return_value={}):
|
|
cfg = load_oidc_config()
|
|
|
|
assert cfg.redirect_base == "http://localhost:8000"
|
|
|
|
def test_load_oidc_config_parses_trusted_endpoint_hosts(self, monkeypatch):
|
|
"""TURNSTONE_OIDC_TRUSTED_ENDPOINT_HOSTS is split, lowercased, and trimmed."""
|
|
monkeypatch.setenv("TURNSTONE_OIDC_ISSUER", "https://auth.example.com")
|
|
monkeypatch.setenv("TURNSTONE_OIDC_CLIENT_ID", "cid")
|
|
monkeypatch.setenv("TURNSTONE_OIDC_CLIENT_SECRET", "csecret")
|
|
monkeypatch.setenv(
|
|
"TURNSTONE_OIDC_TRUSTED_ENDPOINT_HOSTS",
|
|
"Foo.Example.com, BAR.example.com ,, ",
|
|
)
|
|
|
|
with patch("turnstone.core.config.load_config", return_value={}):
|
|
cfg = load_oidc_config()
|
|
|
|
assert cfg.trusted_endpoint_hosts == ("foo.example.com", "bar.example.com")
|
|
|
|
def test_load_oidc_config_trusted_endpoint_hosts_default_empty(self, monkeypatch):
|
|
"""trusted_endpoint_hosts defaults to () when env var is absent."""
|
|
monkeypatch.setenv("TURNSTONE_OIDC_ISSUER", "https://auth.example.com")
|
|
monkeypatch.setenv("TURNSTONE_OIDC_CLIENT_ID", "cid")
|
|
monkeypatch.setenv("TURNSTONE_OIDC_CLIENT_SECRET", "csecret")
|
|
monkeypatch.delenv("TURNSTONE_OIDC_TRUSTED_ENDPOINT_HOSTS", raising=False)
|
|
|
|
with patch("turnstone.core.config.load_config", return_value={}):
|
|
cfg = load_oidc_config()
|
|
|
|
assert cfg.trusted_endpoint_hosts == ()
|
|
|
|
def test_load_oidc_config_trusted_endpoint_hosts_from_toml_list(self, monkeypatch):
|
|
"""config.toml may provide trusted_endpoint_hosts as a list."""
|
|
monkeypatch.setenv("TURNSTONE_OIDC_ISSUER", "https://auth.example.com")
|
|
monkeypatch.setenv("TURNSTONE_OIDC_CLIENT_ID", "cid")
|
|
monkeypatch.setenv("TURNSTONE_OIDC_CLIENT_SECRET", "csecret")
|
|
monkeypatch.delenv("TURNSTONE_OIDC_TRUSTED_ENDPOINT_HOSTS", raising=False)
|
|
|
|
with patch(
|
|
"turnstone.core.config.load_config",
|
|
return_value={"trusted_endpoint_hosts": ["FOO.example.com", "bar.example.com"]},
|
|
):
|
|
cfg = load_oidc_config()
|
|
|
|
assert cfg.trusted_endpoint_hosts == ("foo.example.com", "bar.example.com")
|
|
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# SSRF Validation
|
|
# ---------------------------------------------------------------------------
|
|
|
|
|
|
class TestValidateIssuerURL:
|
|
"""Tests for ``validate_issuer_url`` SSRF protection."""
|
|
|
|
def test_valid_https_url(self):
|
|
"""Public HTTPS issuer URL passes validation."""
|
|
# Should not raise -- mock DNS to return a public IP.
|
|
with patch(
|
|
"socket.getaddrinfo",
|
|
return_value=[
|
|
(2, 1, 6, "", ("93.184.216.34", 0)),
|
|
],
|
|
):
|
|
validate_issuer_url("https://idp.example.com")
|
|
|
|
def test_private_address_hint_mentions_opt_in(self):
|
|
"""The rejection message points the operator at allow_private_network."""
|
|
with (
|
|
patch("socket.getaddrinfo", return_value=[(2, 1, 6, "", ("10.0.0.5", 0))]),
|
|
pytest.raises(OIDCError, match="allow_private_network"),
|
|
):
|
|
validate_issuer_url("https://auth.internal.example")
|
|
|
|
def test_allow_private_accepts_private_issuer(self):
|
|
"""The opt-in accepts an issuer resolving to RFC 1918 space."""
|
|
with patch("socket.getaddrinfo", return_value=[(2, 1, 6, "", ("10.0.0.5", 0))]):
|
|
validate_issuer_url("https://auth.internal.example", allow_private=True)
|
|
|
|
def test_allow_private_still_rejects_link_local(self):
|
|
"""Link-local (cloud metadata) is refused even with the opt-in."""
|
|
with (
|
|
patch("socket.getaddrinfo", return_value=[(2, 1, 6, "", ("169.254.169.254", 0))]),
|
|
pytest.raises(OIDCError, match="refused even with private"),
|
|
):
|
|
validate_issuer_url("https://md.internal.example", allow_private=True)
|
|
|
|
def test_rejects_http_non_localhost(self):
|
|
"""HTTP is rejected for non-localhost hosts."""
|
|
with pytest.raises(OIDCError, match="must use HTTPS"):
|
|
validate_issuer_url("http://idp.example.com")
|
|
|
|
def test_allows_http_localhost(self):
|
|
"""HTTP is allowed for localhost (development)."""
|
|
with patch(
|
|
"socket.getaddrinfo",
|
|
return_value=[
|
|
(2, 1, 6, "", ("127.0.0.1", 0)),
|
|
],
|
|
):
|
|
validate_issuer_url("http://localhost:8080")
|
|
|
|
def test_allows_http_localhost_subdomain(self):
|
|
"""HTTP is allowed for *.localhost subdomains."""
|
|
with patch(
|
|
"socket.getaddrinfo",
|
|
return_value=[
|
|
(2, 1, 6, "", ("127.0.0.1", 0)),
|
|
],
|
|
):
|
|
validate_issuer_url("http://keycloak.localhost:8080")
|
|
|
|
def test_rejects_embedded_credentials(self):
|
|
"""URLs with userinfo (user:pass@host) are rejected."""
|
|
with pytest.raises(OIDCError, match="embedded credentials"):
|
|
validate_issuer_url("https://admin:secret@idp.example.com")
|
|
|
|
def test_rejects_username_only(self):
|
|
"""URLs with just a username are rejected."""
|
|
with pytest.raises(OIDCError, match="embedded credentials"):
|
|
validate_issuer_url("https://admin@idp.example.com")
|
|
|
|
def test_rejects_no_hostname(self):
|
|
"""URLs without a hostname are rejected."""
|
|
with pytest.raises(OIDCError, match="no hostname"):
|
|
validate_issuer_url("https://")
|
|
|
|
def test_rejects_private_10_range(self):
|
|
"""Hostnames resolving to 10.x.x.x are rejected."""
|
|
with (
|
|
patch("socket.getaddrinfo", return_value=[(2, 1, 6, "", ("10.0.0.1", 0))]),
|
|
pytest.raises(OIDCError, match="non-public address.*10.0.0.1"),
|
|
):
|
|
validate_issuer_url("https://internal.corp.example.com")
|
|
|
|
def test_rejects_private_172_range(self):
|
|
"""Hostnames resolving to 172.16-31.x.x are rejected."""
|
|
with (
|
|
patch("socket.getaddrinfo", return_value=[(2, 1, 6, "", ("172.16.0.1", 0))]),
|
|
pytest.raises(OIDCError, match="non-public address.*172.16.0.1"),
|
|
):
|
|
validate_issuer_url("https://internal.corp.example.com")
|
|
|
|
def test_rejects_private_192_168_range(self):
|
|
"""Hostnames resolving to 192.168.x.x are rejected."""
|
|
with (
|
|
patch("socket.getaddrinfo", return_value=[(2, 1, 6, "", ("192.168.1.1", 0))]),
|
|
pytest.raises(OIDCError, match="non-public address.*192.168.1.1"),
|
|
):
|
|
validate_issuer_url("https://internal.corp.example.com")
|
|
|
|
def test_rejects_loopback_127(self):
|
|
"""Hostnames resolving to 127.x.x.x are rejected (non-localhost host)."""
|
|
with (
|
|
patch("socket.getaddrinfo", return_value=[(2, 1, 6, "", ("127.0.0.1", 0))]),
|
|
pytest.raises(OIDCError, match="non-public address.*127.0.0.1"),
|
|
):
|
|
validate_issuer_url("https://evil.example.com")
|
|
|
|
def test_rejects_ipv6_loopback(self):
|
|
"""Hostnames resolving to ::1 are rejected (non-localhost host)."""
|
|
with (
|
|
patch("socket.getaddrinfo", return_value=[(10, 1, 6, "", ("::1", 0, 0, 0))]),
|
|
pytest.raises(OIDCError, match="non-public address.*::1"),
|
|
):
|
|
validate_issuer_url("https://evil.example.com")
|
|
|
|
def test_rejects_ipv6_private(self):
|
|
"""Hostnames resolving to fc00::/7 are rejected."""
|
|
with (
|
|
patch("socket.getaddrinfo", return_value=[(10, 1, 6, "", ("fd00::1", 0, 0, 0))]),
|
|
pytest.raises(OIDCError, match="non-public address.*fd00::1"),
|
|
):
|
|
validate_issuer_url("https://evil.example.com")
|
|
|
|
def test_rejects_link_local(self):
|
|
"""Hostnames resolving to link-local addresses are rejected.
|
|
|
|
Refused as link-local rather than as merely non-public, and WITHOUT the
|
|
allow_private_network hint: that opt-in never admits link-local, so
|
|
offering it would send the operator to a dead end.
|
|
"""
|
|
with (
|
|
patch("socket.getaddrinfo", return_value=[(2, 1, 6, "", ("169.254.169.254", 0))]),
|
|
pytest.raises(OIDCError, match="link-local.*169.254.169.254") as exc_info,
|
|
):
|
|
validate_issuer_url("https://metadata.internal")
|
|
assert "allow_private_network" not in str(exc_info.value)
|
|
|
|
def test_rejects_unresolvable_hostname(self):
|
|
"""DNS resolution failure is rejected."""
|
|
import socket as _socket
|
|
|
|
with (
|
|
patch("socket.getaddrinfo", side_effect=_socket.gaierror("not found")),
|
|
pytest.raises(OIDCError, match="cannot be resolved"),
|
|
):
|
|
validate_issuer_url("https://nonexistent.invalid")
|
|
|
|
def test_rejects_mixed_addresses(self):
|
|
"""If any resolved address is private, the URL is rejected."""
|
|
with (
|
|
patch(
|
|
"socket.getaddrinfo",
|
|
return_value=[
|
|
(2, 1, 6, "", ("93.184.216.34", 0)),
|
|
(2, 1, 6, "", ("10.0.0.1", 0)),
|
|
],
|
|
),
|
|
pytest.raises(OIDCError, match="non-public address.*10.0.0.1"),
|
|
):
|
|
validate_issuer_url("https://dual-homed.example.com")
|
|
|
|
def test_discover_rejects_ssrf(self):
|
|
"""discover_oidc returns enabled=False when issuer URL fails SSRF check."""
|
|
config = _make_config(
|
|
issuer="http://10.0.0.1:8080",
|
|
authorization_endpoint="",
|
|
token_endpoint="",
|
|
userinfo_endpoint="",
|
|
jwks_uri="",
|
|
)
|
|
|
|
async def _run():
|
|
result = await discover_oidc(config)
|
|
assert result.enabled is False
|
|
|
|
asyncio.run(_run())
|
|
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# Discovered Endpoint Validation
|
|
# ---------------------------------------------------------------------------
|
|
|
|
|
|
class TestValidateDiscoveredEndpoint:
|
|
"""Tests for ``validate_discovered_endpoint`` SSRF + same-origin protection."""
|
|
|
|
_PUBLIC_ADDR = [(2, 1, 6, "", ("93.184.216.34", 0))]
|
|
_PRIVATE_ADDR = [(2, 1, 6, "", ("169.254.169.254", 0))]
|
|
_LOOPBACK_ADDR = [(2, 1, 6, "", ("127.0.0.1", 0))]
|
|
|
|
@staticmethod
|
|
def _issuer(url: str = "https://idp.example.com") -> urllib.parse.ParseResult:
|
|
return urllib.parse.urlparse(url)
|
|
|
|
def test_valid_same_origin(self):
|
|
"""Same scheme/host/port as issuer passes."""
|
|
with patch("socket.getaddrinfo", return_value=self._PUBLIC_ADDR):
|
|
validate_discovered_endpoint(
|
|
"https://idp.example.com/token",
|
|
self._issuer(),
|
|
allow_http=False,
|
|
trusted_endpoint_hosts=frozenset(),
|
|
)
|
|
|
|
def test_private_endpoint_hint_mentions_opt_in(self):
|
|
"""A discovered endpoint resolving private carries the opt-in hint
|
|
just like the issuer does — the remediation is the same knob."""
|
|
with (
|
|
patch("socket.getaddrinfo", return_value=[(2, 1, 6, "", ("10.0.0.5", 0))]),
|
|
pytest.raises(OIDCError, match="allow_private_network"),
|
|
):
|
|
validate_discovered_endpoint(
|
|
"https://idp.example.com/token",
|
|
self._issuer(),
|
|
allow_http=False,
|
|
trusted_endpoint_hosts=frozenset(),
|
|
)
|
|
|
|
def test_rejects_http_when_issuer_is_https(self):
|
|
"""http:// discovered endpoint rejected when issuer is https://."""
|
|
with (
|
|
patch("socket.getaddrinfo", return_value=self._PUBLIC_ADDR),
|
|
pytest.raises(OIDCError, match="must use HTTPS"),
|
|
):
|
|
validate_discovered_endpoint(
|
|
"http://idp.example.com/token",
|
|
self._issuer(),
|
|
allow_http=False,
|
|
trusted_endpoint_hosts=frozenset(),
|
|
)
|
|
|
|
def test_allows_http_when_issuer_is_localhost(self):
|
|
"""http:// discovered endpoint allowed in dev when issuer is http://localhost."""
|
|
with patch("socket.getaddrinfo", return_value=self._LOOPBACK_ADDR):
|
|
validate_discovered_endpoint(
|
|
"http://localhost:8080/token",
|
|
self._issuer("http://localhost:8080"),
|
|
allow_http=True,
|
|
trusted_endpoint_hosts=frozenset(),
|
|
)
|
|
|
|
def test_rejects_link_local_ip(self):
|
|
"""Endpoint resolving to a link-local IP is rejected, with no dead-end hint."""
|
|
with (
|
|
patch("socket.getaddrinfo", return_value=self._PRIVATE_ADDR),
|
|
pytest.raises(OIDCError, match="link-local.*169.254.169.254") as exc_info,
|
|
):
|
|
validate_discovered_endpoint(
|
|
"https://idp.example.com/token",
|
|
self._issuer(),
|
|
allow_http=False,
|
|
trusted_endpoint_hosts=frozenset(),
|
|
)
|
|
assert "allow_private_network" not in str(exc_info.value)
|
|
|
|
def test_rejects_private_ip(self):
|
|
"""Endpoint resolving to a genuinely private IP is rejected — and IS hinted.
|
|
|
|
``_PRIVATE_ADDR`` is 169.254.169.254, which is link-local rather than
|
|
private, so this case covers what the name promises: an RFC 1918 address
|
|
the operator CAN reach by enabling the opt-in, and is told so.
|
|
"""
|
|
with (
|
|
patch("socket.getaddrinfo", return_value=[(2, 1, 6, "", ("10.0.0.7", 0))]),
|
|
pytest.raises(OIDCError, match="non-public address.*10.0.0.7") as exc_info,
|
|
):
|
|
validate_discovered_endpoint(
|
|
"https://idp.example.com/token",
|
|
self._issuer(),
|
|
allow_http=False,
|
|
trusted_endpoint_hosts=frozenset(),
|
|
)
|
|
assert "allow_private_network" in str(exc_info.value)
|
|
|
|
def test_rejects_embedded_credentials(self):
|
|
"""Endpoint with userinfo (user:pass@host) rejected."""
|
|
with pytest.raises(OIDCError, match="embedded credentials"):
|
|
validate_discovered_endpoint(
|
|
"https://user:pass@idp.example.com/token",
|
|
self._issuer(),
|
|
allow_http=False,
|
|
trusted_endpoint_hosts=frozenset(),
|
|
)
|
|
|
|
def test_rejects_different_host(self):
|
|
"""Endpoint on a different host than the issuer is rejected."""
|
|
with (
|
|
patch("socket.getaddrinfo", return_value=self._PUBLIC_ADDR),
|
|
pytest.raises(OIDCError, match="not trusted"),
|
|
):
|
|
validate_discovered_endpoint(
|
|
"https://attacker.com/token",
|
|
self._issuer(),
|
|
allow_http=False,
|
|
trusted_endpoint_hosts=frozenset(),
|
|
)
|
|
|
|
def test_rejects_subdomain(self):
|
|
"""Sibling-subdomain endpoint is rejected (strict equality)."""
|
|
with (
|
|
patch("socket.getaddrinfo", return_value=self._PUBLIC_ADDR),
|
|
pytest.raises(OIDCError, match="not trusted"),
|
|
):
|
|
validate_discovered_endpoint(
|
|
"https://login.idp.example.com/token",
|
|
self._issuer(),
|
|
allow_http=False,
|
|
trusted_endpoint_hosts=frozenset(),
|
|
)
|
|
|
|
def test_rejects_different_port(self):
|
|
"""Endpoint on a different port than the issuer is rejected."""
|
|
with (
|
|
patch("socket.getaddrinfo", return_value=self._PUBLIC_ADDR),
|
|
pytest.raises(OIDCError, match="port.*does not match issuer"),
|
|
):
|
|
validate_discovered_endpoint(
|
|
"https://idp.example.com:9443/token",
|
|
self._issuer(),
|
|
allow_http=False,
|
|
trusted_endpoint_hosts=frozenset(),
|
|
)
|
|
|
|
def test_rejects_different_scheme(self):
|
|
"""Endpoint scheme must match the issuer's scheme."""
|
|
with (
|
|
patch("socket.getaddrinfo", return_value=self._LOOPBACK_ADDR),
|
|
pytest.raises(OIDCError, match="scheme.*does not match issuer"),
|
|
):
|
|
validate_discovered_endpoint(
|
|
"https://localhost:8080/token",
|
|
self._issuer("http://localhost:8080"),
|
|
allow_http=True,
|
|
trusted_endpoint_hosts=frozenset(),
|
|
)
|
|
|
|
def test_accepts_google_known_endpoints(self):
|
|
"""Issuer accounts.google.com accepts the four well-known multi-origin hosts."""
|
|
google_endpoints = (
|
|
"https://accounts.google.com/o/oauth2/v2/auth",
|
|
"https://oauth2.googleapis.com/token",
|
|
"https://www.googleapis.com/oauth2/v3/certs",
|
|
"https://openidconnect.googleapis.com/v1/userinfo",
|
|
)
|
|
with patch("socket.getaddrinfo", return_value=self._PUBLIC_ADDR):
|
|
for endpoint in google_endpoints:
|
|
validate_discovered_endpoint(
|
|
endpoint,
|
|
self._issuer("https://accounts.google.com"),
|
|
allow_http=False,
|
|
trusted_endpoint_hosts=frozenset(),
|
|
)
|
|
|
|
def test_rejects_unknown_host_for_known_issuer(self):
|
|
"""Even with a known issuer, foreign endpoints outside the allow-map are rejected."""
|
|
with (
|
|
patch("socket.getaddrinfo", return_value=self._PUBLIC_ADDR),
|
|
pytest.raises(OIDCError, match="not trusted"),
|
|
):
|
|
validate_discovered_endpoint(
|
|
"https://attacker.com/token",
|
|
self._issuer("https://accounts.google.com"),
|
|
allow_http=False,
|
|
trusted_endpoint_hosts=frozenset(),
|
|
)
|
|
|
|
def test_accepts_operator_trusted_endpoint_host(self):
|
|
"""Operator-supplied trusted_endpoint_hosts permits cross-host endpoints."""
|
|
with patch("socket.getaddrinfo", return_value=self._PUBLIC_ADDR):
|
|
validate_discovered_endpoint(
|
|
"https://idp-token.example.net/token",
|
|
self._issuer("https://idp.example.com"),
|
|
allow_http=False,
|
|
trusted_endpoint_hosts=frozenset({"idp-token.example.net"}),
|
|
)
|
|
|
|
def test_accepts_explicit_default_port_match(self):
|
|
"""Issuer omits :443; endpoint includes :443 explicitly -> still same origin."""
|
|
with patch("socket.getaddrinfo", return_value=self._PUBLIC_ADDR):
|
|
validate_discovered_endpoint(
|
|
"https://idp.example.com:443/token",
|
|
self._issuer("https://idp.example.com"),
|
|
allow_http=False,
|
|
trusted_endpoint_hosts=frozenset(),
|
|
)
|
|
|
|
def test_accepts_implicit_default_port_match_reverse(self):
|
|
"""Issuer includes :443; endpoint omits the port -> still same origin."""
|
|
with patch("socket.getaddrinfo", return_value=self._PUBLIC_ADDR):
|
|
validate_discovered_endpoint(
|
|
"https://idp.example.com/token",
|
|
self._issuer("https://idp.example.com:443"),
|
|
allow_http=False,
|
|
trusted_endpoint_hosts=frozenset(),
|
|
)
|
|
|
|
def test_discover_rejects_endpoint_on_other_host(self):
|
|
"""discover_oidc returns enabled=False when token_endpoint targets a foreign host."""
|
|
config = _make_config(
|
|
authorization_endpoint="",
|
|
token_endpoint="",
|
|
userinfo_endpoint="",
|
|
jwks_uri="",
|
|
)
|
|
discovery_doc = {
|
|
"authorization_endpoint": "https://idp.example.com/authorize",
|
|
"token_endpoint": "https://attacker.com/token",
|
|
"userinfo_endpoint": "https://idp.example.com/userinfo",
|
|
"jwks_uri": "https://idp.example.com/.well-known/jwks.json",
|
|
}
|
|
mock_response = MagicMock()
|
|
mock_response.json.return_value = discovery_doc
|
|
mock_response.raise_for_status = MagicMock()
|
|
|
|
async def _run():
|
|
client = _mock_async_client(lambda url: _async_return(mock_response))
|
|
with (
|
|
patch("socket.getaddrinfo", return_value=self._PUBLIC_ADDR),
|
|
patch("httpx.AsyncClient", return_value=client),
|
|
):
|
|
result = await discover_oidc(config)
|
|
assert result.enabled is False
|
|
|
|
asyncio.run(_run())
|
|
|
|
def test_discover_rejects_foreign_userinfo(self):
|
|
"""A foreign userinfo_endpoint is rejected even when other endpoints are same-origin.
|
|
|
|
The required-endpoint loop runs first, so this test pins the userinfo branch's
|
|
reject path and proves it executes.
|
|
"""
|
|
config = _make_config(
|
|
authorization_endpoint="",
|
|
token_endpoint="",
|
|
userinfo_endpoint="",
|
|
jwks_uri="",
|
|
)
|
|
discovery_doc = {
|
|
"authorization_endpoint": "https://idp.example.com/authorize",
|
|
"token_endpoint": "https://idp.example.com/token",
|
|
"userinfo_endpoint": "https://attacker.com/userinfo",
|
|
"jwks_uri": "https://idp.example.com/.well-known/jwks.json",
|
|
}
|
|
mock_response = MagicMock()
|
|
mock_response.json.return_value = discovery_doc
|
|
mock_response.raise_for_status = MagicMock()
|
|
|
|
async def _run():
|
|
client = _mock_async_client(lambda url: _async_return(mock_response))
|
|
with (
|
|
patch("socket.getaddrinfo", return_value=self._PUBLIC_ADDR),
|
|
patch("httpx.AsyncClient", return_value=client),
|
|
):
|
|
result = await discover_oidc(config)
|
|
assert result.enabled is False
|
|
|
|
asyncio.run(_run())
|
|
|
|
def test_discover_accepts_google_multi_origin(self):
|
|
"""discover_oidc accepts Google's legitimate multi-origin discovery doc."""
|
|
config = _make_config(
|
|
issuer="https://accounts.google.com",
|
|
authorization_endpoint="",
|
|
token_endpoint="",
|
|
userinfo_endpoint="",
|
|
jwks_uri="",
|
|
)
|
|
discovery_doc = {
|
|
"authorization_endpoint": "https://accounts.google.com/o/oauth2/v2/auth",
|
|
"token_endpoint": "https://oauth2.googleapis.com/token",
|
|
"userinfo_endpoint": "https://openidconnect.googleapis.com/v1/userinfo",
|
|
"jwks_uri": "https://www.googleapis.com/oauth2/v3/certs",
|
|
}
|
|
mock_response = MagicMock()
|
|
mock_response.json.return_value = discovery_doc
|
|
mock_response.raise_for_status = MagicMock()
|
|
|
|
async def _run():
|
|
client = _mock_async_client(lambda url: _async_return(mock_response))
|
|
with (
|
|
patch("socket.getaddrinfo", return_value=self._PUBLIC_ADDR),
|
|
patch("httpx.AsyncClient", return_value=client),
|
|
):
|
|
result = await discover_oidc(config)
|
|
assert result.enabled is True
|
|
assert result.token_endpoint == "https://oauth2.googleapis.com/token"
|
|
assert result.jwks_uri == "https://www.googleapis.com/oauth2/v3/certs"
|
|
|
|
asyncio.run(_run())
|
|
|
|
def test_discover_accepts_entra_userinfo_on_graph(self):
|
|
"""discover_oidc accepts Entra's cross-host userinfo on graph.microsoft.com
|
|
via the built-in allow-list, so Azure AD OIDC works out of the box with no
|
|
trusted_endpoint_hosts override."""
|
|
tenant = "11111111-1111-1111-1111-111111111111"
|
|
config = _make_config(
|
|
issuer=f"https://login.microsoftonline.com/{tenant}/v2.0",
|
|
authorization_endpoint="",
|
|
token_endpoint="",
|
|
userinfo_endpoint="",
|
|
jwks_uri="",
|
|
)
|
|
discovery_doc = {
|
|
"authorization_endpoint": f"https://login.microsoftonline.com/{tenant}/oauth2/v2.0/authorize",
|
|
"token_endpoint": f"https://login.microsoftonline.com/{tenant}/oauth2/v2.0/token",
|
|
"userinfo_endpoint": "https://graph.microsoft.com/oidc/userinfo",
|
|
"jwks_uri": f"https://login.microsoftonline.com/{tenant}/discovery/v2.0/keys",
|
|
}
|
|
mock_response = MagicMock()
|
|
mock_response.json.return_value = discovery_doc
|
|
mock_response.raise_for_status = MagicMock()
|
|
|
|
async def _run():
|
|
client = _mock_async_client(lambda url: _async_return(mock_response))
|
|
with (
|
|
patch("socket.getaddrinfo", return_value=self._PUBLIC_ADDR),
|
|
patch("httpx.AsyncClient", return_value=client),
|
|
):
|
|
result = await discover_oidc(config)
|
|
assert result.enabled is True
|
|
assert result.userinfo_endpoint == "https://graph.microsoft.com/oidc/userinfo"
|
|
|
|
asyncio.run(_run())
|
|
|
|
def test_discover_rejects_http_endpoint(self):
|
|
"""discover_oidc returns enabled=False when an endpoint is http:// in prod."""
|
|
config = _make_config(
|
|
authorization_endpoint="",
|
|
token_endpoint="",
|
|
userinfo_endpoint="",
|
|
jwks_uri="",
|
|
)
|
|
discovery_doc = {
|
|
"authorization_endpoint": "https://idp.example.com/authorize",
|
|
"token_endpoint": "http://idp.example.com/token",
|
|
"userinfo_endpoint": "https://idp.example.com/userinfo",
|
|
"jwks_uri": "https://idp.example.com/.well-known/jwks.json",
|
|
}
|
|
mock_response = MagicMock()
|
|
mock_response.json.return_value = discovery_doc
|
|
mock_response.raise_for_status = MagicMock()
|
|
|
|
async def _run():
|
|
client = _mock_async_client(lambda url: _async_return(mock_response))
|
|
with (
|
|
patch("socket.getaddrinfo", return_value=self._PUBLIC_ADDR),
|
|
patch("httpx.AsyncClient", return_value=client),
|
|
):
|
|
result = await discover_oidc(config)
|
|
assert result.enabled is False
|
|
|
|
asyncio.run(_run())
|
|
|
|
def test_discover_rejects_endpoint_resolving_to_private_ip(self):
|
|
"""discover_oidc returns enabled=False when an endpoint host resolves privately.
|
|
|
|
DNS may legitimately rotate between the issuer check and the per-endpoint
|
|
re-resolution; defence-in-depth requires we re-validate each discovered URL.
|
|
"""
|
|
config = _make_config(
|
|
authorization_endpoint="",
|
|
token_endpoint="",
|
|
userinfo_endpoint="",
|
|
jwks_uri="",
|
|
)
|
|
discovery_doc = {
|
|
"authorization_endpoint": "https://idp.example.com/authorize",
|
|
"token_endpoint": "https://idp.example.com/token",
|
|
"userinfo_endpoint": "https://idp.example.com/userinfo",
|
|
"jwks_uri": "https://idp.example.com/.well-known/jwks.json",
|
|
}
|
|
mock_response = MagicMock()
|
|
mock_response.json.return_value = discovery_doc
|
|
mock_response.raise_for_status = MagicMock()
|
|
|
|
# Issuer validation passes (public), then DNS rotates so each discovered
|
|
# endpoint resolves to a link-local address.
|
|
results = iter(
|
|
[
|
|
self._PUBLIC_ADDR,
|
|
self._PRIVATE_ADDR,
|
|
self._PRIVATE_ADDR,
|
|
self._PRIVATE_ADDR,
|
|
self._PRIVATE_ADDR,
|
|
]
|
|
)
|
|
|
|
def _resolve(*_args, **_kwargs):
|
|
return next(results)
|
|
|
|
async def _run():
|
|
client = _mock_async_client(lambda url: _async_return(mock_response))
|
|
with (
|
|
patch("socket.getaddrinfo", side_effect=_resolve),
|
|
patch("httpx.AsyncClient", return_value=client),
|
|
):
|
|
result = await discover_oidc(config)
|
|
assert result.enabled is False
|
|
|
|
asyncio.run(_run())
|
|
|
|
def test_discover_accepts_localhost_http_flow(self):
|
|
"""Localhost issuer with http:// endpoints is accepted in dev mode."""
|
|
config = _make_config(
|
|
issuer="http://localhost:8080",
|
|
authorization_endpoint="",
|
|
token_endpoint="",
|
|
userinfo_endpoint="",
|
|
jwks_uri="",
|
|
)
|
|
discovery_doc = {
|
|
"authorization_endpoint": "http://localhost:8080/authorize",
|
|
"token_endpoint": "http://localhost:8080/token",
|
|
"userinfo_endpoint": "http://localhost:8080/userinfo",
|
|
"jwks_uri": "http://localhost:8080/jwks",
|
|
}
|
|
mock_response = MagicMock()
|
|
mock_response.json.return_value = discovery_doc
|
|
mock_response.raise_for_status = MagicMock()
|
|
|
|
async def _run():
|
|
client = _mock_async_client(lambda url: _async_return(mock_response))
|
|
with (
|
|
patch("socket.getaddrinfo", return_value=self._LOOPBACK_ADDR),
|
|
patch("httpx.AsyncClient", return_value=client),
|
|
):
|
|
result = await discover_oidc(config)
|
|
assert result.enabled is True
|
|
assert result.token_endpoint == "http://localhost:8080/token"
|
|
|
|
asyncio.run(_run())
|
|
|
|
def test_discover_accepts_empty_userinfo(self):
|
|
"""Empty userinfo_endpoint is allowed (skipped) when other endpoints are valid."""
|
|
config = _make_config(
|
|
authorization_endpoint="",
|
|
token_endpoint="",
|
|
userinfo_endpoint="",
|
|
jwks_uri="",
|
|
)
|
|
discovery_doc = {
|
|
"authorization_endpoint": "https://idp.example.com/authorize",
|
|
"token_endpoint": "https://idp.example.com/token",
|
|
"jwks_uri": "https://idp.example.com/.well-known/jwks.json",
|
|
}
|
|
mock_response = MagicMock()
|
|
mock_response.json.return_value = discovery_doc
|
|
mock_response.raise_for_status = MagicMock()
|
|
|
|
async def _run():
|
|
client = _mock_async_client(lambda url: _async_return(mock_response))
|
|
with (
|
|
patch("socket.getaddrinfo", return_value=self._PUBLIC_ADDR),
|
|
patch("httpx.AsyncClient", return_value=client),
|
|
):
|
|
result = await discover_oidc(config)
|
|
assert result.enabled is True
|
|
assert result.userinfo_endpoint == ""
|
|
|
|
asyncio.run(_run())
|
|
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# Redirect URI Builder
|
|
# ---------------------------------------------------------------------------
|
|
|
|
|
|
class TestBuildOIDCRedirectURI:
|
|
"""Tests for ``_build_oidc_redirect_uri`` in auth.py."""
|
|
|
|
def test_build_redirect_uri_uses_redirect_base_only(self):
|
|
"""The redirect URI is built solely from ``redirect_base``."""
|
|
from turnstone.core.auth import _build_oidc_redirect_uri
|
|
|
|
config = _make_config(redirect_base="https://example.com")
|
|
result = _build_oidc_redirect_uri(config)
|
|
assert result == "https://example.com/v1/api/auth/oidc/callback"
|
|
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# PKCE
|
|
# ---------------------------------------------------------------------------
|
|
|
|
|
|
class TestPKCE:
|
|
def test_generate_pkce_verifier_shape(self):
|
|
verifier = generate_pkce_verifier()
|
|
|
|
# Verifier should be URL-safe base64
|
|
assert isinstance(verifier, str)
|
|
assert len(verifier) > 40 # 48 bytes -> ~64 chars
|
|
|
|
def test_pkce_verifier_uniqueness(self):
|
|
"""Each call should produce a unique verifier."""
|
|
v1 = generate_pkce_verifier()
|
|
v2 = generate_pkce_verifier()
|
|
assert v1 != v2
|
|
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# Authorization URL
|
|
# ---------------------------------------------------------------------------
|
|
|
|
|
|
class TestBuildAuthorizeURL:
|
|
def test_build_authorize_url_contains_required_params(self):
|
|
config = _make_config()
|
|
verifier = generate_pkce_verifier()
|
|
url = build_authorize_url(
|
|
config=config,
|
|
redirect_uri="https://app.example.com/callback",
|
|
state="test-state",
|
|
nonce="test-nonce",
|
|
code_verifier=verifier,
|
|
)
|
|
|
|
assert url.startswith("https://idp.example.com/authorize?")
|
|
assert "response_type=code" in url
|
|
assert "client_id=my-client" in url
|
|
assert "redirect_uri=" in url
|
|
assert "scope=openid" in url
|
|
assert "state=test-state" in url
|
|
assert "nonce=test-nonce" in url
|
|
assert "code_challenge=" in url
|
|
assert "code_challenge_method=S256" in url
|
|
|
|
def test_build_authorize_url_pkce(self):
|
|
"""code_challenge in URL should be correct S256 of the verifier."""
|
|
config = _make_config()
|
|
verifier = generate_pkce_verifier()
|
|
|
|
url = build_authorize_url(
|
|
config=config,
|
|
redirect_uri="https://app.example.com/callback",
|
|
state="s",
|
|
nonce="n",
|
|
code_verifier=verifier,
|
|
)
|
|
|
|
# Extract code_challenge from URL
|
|
parsed = urllib.parse.urlparse(url)
|
|
params = urllib.parse.parse_qs(parsed.query)
|
|
actual_challenge = params["code_challenge"][0]
|
|
|
|
# Compute expected challenge
|
|
digest = hashlib.sha256(verifier.encode("ascii")).digest()
|
|
expected = base64.urlsafe_b64encode(digest).rstrip(b"=").decode("ascii")
|
|
assert actual_challenge == expected
|
|
|
|
def test_build_authorize_url_redirect_uri_encoded(self):
|
|
config = _make_config()
|
|
verifier = generate_pkce_verifier()
|
|
redirect = "https://app.example.com/callback?extra=1"
|
|
|
|
url = build_authorize_url(
|
|
config=config,
|
|
redirect_uri=redirect,
|
|
state="s",
|
|
nonce="n",
|
|
code_verifier=verifier,
|
|
)
|
|
|
|
# The redirect_uri should be URL-encoded
|
|
parsed = urllib.parse.urlparse(url)
|
|
params = urllib.parse.parse_qs(parsed.query)
|
|
assert params["redirect_uri"][0] == redirect
|
|
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# ID Token Validation
|
|
# ---------------------------------------------------------------------------
|
|
|
|
|
|
class TestValidateIDToken:
|
|
_FAKE_JWKS = {"keys": [{"kid": "key1", "kty": "RSA", "n": "abc", "e": "AQAB"}]}
|
|
|
|
def test_validate_id_token_nonce_mismatch(self):
|
|
"""Nonce mismatch should raise OIDCError."""
|
|
config = _make_config()
|
|
|
|
mock_pyjwk = MagicMock()
|
|
mock_pyjwk.return_value.key = "fake-key"
|
|
|
|
with (
|
|
patch("jwt.get_unverified_header", return_value={"kid": "key1", "alg": "RS256"}),
|
|
patch("jwt.PyJWK", mock_pyjwk),
|
|
patch("jwt.decode", return_value={"sub": "user1", "nonce": "wrong-nonce"}),
|
|
pytest.raises(OIDCError, match="nonce mismatch"),
|
|
):
|
|
validate_id_token(
|
|
raw_token="fake.jwt.token",
|
|
jwks_data=self._FAKE_JWKS,
|
|
config=config,
|
|
nonce="expected-nonce",
|
|
)
|
|
|
|
def test_validate_id_token_success(self):
|
|
"""Successful validation returns decoded claims."""
|
|
config = _make_config()
|
|
|
|
mock_pyjwk = MagicMock()
|
|
mock_pyjwk.return_value.key = "fake-key"
|
|
|
|
expected_claims = {
|
|
"sub": "user1",
|
|
"email": "user@example.com",
|
|
"nonce": "test-nonce",
|
|
}
|
|
|
|
with (
|
|
patch("jwt.get_unverified_header", return_value={"kid": "key1", "alg": "RS256"}),
|
|
patch("jwt.PyJWK", mock_pyjwk),
|
|
patch("jwt.decode", return_value=expected_claims) as mock_decode,
|
|
):
|
|
claims = validate_id_token(
|
|
raw_token="fake.jwt.token",
|
|
jwks_data=self._FAKE_JWKS,
|
|
config=config,
|
|
nonce="test-nonce",
|
|
)
|
|
|
|
assert claims == expected_claims
|
|
mock_decode.assert_called_once_with(
|
|
"fake.jwt.token",
|
|
"fake-key",
|
|
algorithms=[
|
|
"RS256",
|
|
"RS384",
|
|
"RS512",
|
|
"ES256",
|
|
"ES384",
|
|
"ES512",
|
|
"PS256",
|
|
"PS384",
|
|
"PS512",
|
|
],
|
|
audience="my-client",
|
|
issuer="https://idp.example.com",
|
|
)
|
|
|
|
def test_validate_id_token_kid_not_found(self):
|
|
"""Unknown kid raises OIDCError with descriptive message."""
|
|
config = _make_config()
|
|
jwks_data = {"keys": [{"kid": "other-key", "kty": "RSA"}]}
|
|
|
|
with (
|
|
patch("jwt.get_unverified_header", return_value={"kid": "unknown", "alg": "RS256"}),
|
|
pytest.raises(OIDCError, match="not found in JWKS"),
|
|
):
|
|
validate_id_token(
|
|
raw_token="bad.token",
|
|
jwks_data=jwks_data,
|
|
config=config,
|
|
nonce="n",
|
|
)
|
|
|
|
def test_validate_id_token_raises_keynotfound_for_unknown_kid(self):
|
|
"""Unknown kid raises the OIDCKeyNotFoundError subclass, not generic OIDCError."""
|
|
config = _make_config()
|
|
jwks_data = {"keys": [{"kid": "other-key", "kty": "RSA"}]}
|
|
|
|
with (
|
|
patch("jwt.get_unverified_header", return_value={"kid": "unknown", "alg": "RS256"}),
|
|
pytest.raises(OIDCKeyNotFoundError) as exc_info,
|
|
):
|
|
validate_id_token(
|
|
raw_token="bad.token",
|
|
jwks_data=jwks_data,
|
|
config=config,
|
|
nonce="n",
|
|
)
|
|
assert isinstance(exc_info.value, OIDCError)
|
|
assert "not found in JWKS" in str(exc_info.value)
|
|
|
|
def test_validate_id_token_invalid_jwt(self):
|
|
"""Invalid JWT raises OIDCError."""
|
|
config = _make_config()
|
|
|
|
mock_pyjwk = MagicMock()
|
|
mock_pyjwk.return_value.key = "fake-key"
|
|
|
|
with (
|
|
patch("jwt.get_unverified_header", return_value={"kid": "key1", "alg": "RS256"}),
|
|
patch("jwt.PyJWK", mock_pyjwk := MagicMock(return_value=MagicMock(key="fake-key"))),
|
|
patch("jwt.decode", side_effect=pyjwt.InvalidTokenError("expired")),
|
|
pytest.raises(OIDCError, match="ID token validation failed"),
|
|
):
|
|
validate_id_token(
|
|
raw_token="expired.token",
|
|
jwks_data=self._FAKE_JWKS,
|
|
config=config,
|
|
nonce="n",
|
|
)
|
|
|
|
def test_validate_id_token_invalid_header(self):
|
|
"""Malformed token header raises OIDCError."""
|
|
config = _make_config()
|
|
|
|
with (
|
|
patch("jwt.get_unverified_header", side_effect=pyjwt.DecodeError("bad header")),
|
|
pytest.raises(OIDCError, match="Invalid ID token header"),
|
|
):
|
|
validate_id_token(
|
|
raw_token="garbage",
|
|
jwks_data=self._FAKE_JWKS,
|
|
config=config,
|
|
nonce="n",
|
|
)
|
|
|
|
def test_validate_id_token_retry_after_kid_rotation(self):
|
|
"""Sign with a fresh key, miss old JWKS, succeed against rotated JWKS.
|
|
|
|
Drives the kid-not-found -> JWKS rotation -> retry path at the
|
|
``validate_id_token`` unit level. Uses real RSA signing so the
|
|
decoded claims come back through pyjwt rather than from a mock.
|
|
"""
|
|
from cryptography.hazmat.primitives import serialization
|
|
from cryptography.hazmat.primitives.asymmetric import rsa
|
|
from jwt.algorithms import RSAAlgorithm
|
|
|
|
key = rsa.generate_private_key(public_exponent=65537, key_size=2048)
|
|
priv_pem = key.private_bytes(
|
|
serialization.Encoding.PEM,
|
|
serialization.PrivateFormat.PKCS8,
|
|
serialization.NoEncryption(),
|
|
)
|
|
|
|
new_jwk = RSAAlgorithm.to_jwk(key.public_key(), as_dict=True)
|
|
new_jwk["kid"] = "new-key"
|
|
new_jwk["alg"] = "RS256"
|
|
|
|
config = _make_config(issuer="https://idp.example.com")
|
|
token = pyjwt.encode(
|
|
{
|
|
"sub": "user1",
|
|
"aud": config.client_id,
|
|
"iss": config.issuer,
|
|
"nonce": "test-nonce",
|
|
},
|
|
priv_pem,
|
|
algorithm="RS256",
|
|
headers={"kid": "new-key"},
|
|
)
|
|
|
|
# Stale JWKS with an unrelated old key -> kid lookup fails.
|
|
old_jwks = {
|
|
"keys": [{"kid": "old-key", "kty": "RSA", "n": "abc", "e": "AQAB"}],
|
|
}
|
|
with pytest.raises(OIDCKeyNotFoundError):
|
|
validate_id_token(token, old_jwks, config, "test-nonce")
|
|
|
|
# Rotated JWKS includes the new key -> validation succeeds.
|
|
new_jwks = {"keys": [new_jwk]}
|
|
claims = validate_id_token(token, new_jwks, config, "test-nonce")
|
|
assert claims["sub"] == "user1"
|
|
assert claims["nonce"] == "test-nonce"
|
|
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# Token Exchange
|
|
# ---------------------------------------------------------------------------
|
|
|
|
|
|
def _mock_async_post_client(mock_post):
|
|
"""Build a patched httpx.AsyncClient context manager that supports POST."""
|
|
|
|
class _AsyncCtx:
|
|
async def __aenter__(self):
|
|
return self
|
|
|
|
async def __aexit__(self, *args):
|
|
pass
|
|
|
|
async def post(self, url, data=None):
|
|
return await mock_post(url, data)
|
|
|
|
return _AsyncCtx()
|
|
|
|
|
|
class TestExchangeCode:
|
|
def test_exchange_code_rejects_non_dict_body(self):
|
|
"""A 200 response whose JSON body isn't a JSON object must be rejected."""
|
|
config = _make_config()
|
|
|
|
mock_response = MagicMock()
|
|
mock_response.status_code = 200
|
|
mock_response.json.return_value = [1, 2, 3]
|
|
|
|
async def _post(_url, _data):
|
|
return mock_response
|
|
|
|
async def _run():
|
|
client = _mock_async_post_client(_post)
|
|
with (
|
|
patch("httpx.AsyncClient", return_value=client),
|
|
pytest.raises(OIDCError, match="non-dict body"),
|
|
):
|
|
await exchange_code(config, "code", "https://app.example.com/cb", "verifier")
|
|
|
|
asyncio.run(_run())
|
|
|
|
def test_exchange_code_error_body_is_sanitized(self):
|
|
"""A 4xx response body containing CR/LF must not appear raw in the error message."""
|
|
config = _make_config()
|
|
|
|
mock_response = MagicMock()
|
|
mock_response.status_code = 400
|
|
mock_response.text = "evil\nlog injection\rline"
|
|
|
|
async def _post(_url, _data):
|
|
return mock_response
|
|
|
|
async def _run():
|
|
client = _mock_async_post_client(_post)
|
|
with (
|
|
patch("httpx.AsyncClient", return_value=client),
|
|
pytest.raises(OIDCError) as exc_info,
|
|
):
|
|
await exchange_code(config, "code", "https://app.example.com/cb", "verifier")
|
|
|
|
msg = str(exc_info.value)
|
|
assert "\n" not in msg
|
|
assert "\r" not in msg
|
|
assert "evil" in msg
|
|
assert "log injection" in msg
|
|
|
|
asyncio.run(_run())
|
|
|
|
def test_exchange_code_4xx_status_in_error(self):
|
|
"""A 4xx response surfaces the status code in the OIDCError message."""
|
|
config = _make_config()
|
|
|
|
mock_response = MagicMock()
|
|
mock_response.status_code = 400
|
|
mock_response.text = "invalid_grant"
|
|
|
|
async def _post(_url, _data):
|
|
return mock_response
|
|
|
|
async def _run():
|
|
client = _mock_async_post_client(_post)
|
|
with (
|
|
patch("httpx.AsyncClient", return_value=client),
|
|
pytest.raises(OIDCError, match="returned 400"),
|
|
):
|
|
await exchange_code(config, "code", "https://app.example.com/cb", "verifier")
|
|
|
|
asyncio.run(_run())
|
|
|
|
def test_exchange_code_5xx_status_in_error(self):
|
|
"""A 5xx response surfaces the status code in the OIDCError message."""
|
|
config = _make_config()
|
|
|
|
mock_response = MagicMock()
|
|
mock_response.status_code = 500
|
|
mock_response.text = "internal error"
|
|
|
|
async def _post(_url, _data):
|
|
return mock_response
|
|
|
|
async def _run():
|
|
client = _mock_async_post_client(_post)
|
|
with (
|
|
patch("httpx.AsyncClient", return_value=client),
|
|
pytest.raises(OIDCError, match="returned 500"),
|
|
):
|
|
await exchange_code(config, "code", "https://app.example.com/cb", "verifier")
|
|
|
|
asyncio.run(_run())
|
|
|
|
def test_exchange_code_network_error_wraps_to_oidc_error(self):
|
|
"""``httpx.RequestError`` from the transport surfaces as ``OIDCError``."""
|
|
config = _make_config()
|
|
|
|
async def _post(_url, _data):
|
|
raise httpx.ConnectError("connection refused")
|
|
|
|
async def _run():
|
|
client = _mock_async_post_client(_post)
|
|
with (
|
|
patch("httpx.AsyncClient", return_value=client),
|
|
pytest.raises(OIDCError, match="Token exchange request failed"),
|
|
):
|
|
await exchange_code(config, "code", "https://app.example.com/cb", "verifier")
|
|
|
|
asyncio.run(_run())
|
|
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# JWKS Fetch
|
|
# ---------------------------------------------------------------------------
|
|
|
|
|
|
def _mock_async_get_client(mock_get):
|
|
"""Build a patched httpx.AsyncClient context manager that supports GET-with-timeout."""
|
|
|
|
class _AsyncCtx:
|
|
async def __aenter__(self):
|
|
return self
|
|
|
|
async def __aexit__(self, *args):
|
|
pass
|
|
|
|
async def get(self, url, timeout=None):
|
|
return await mock_get(url)
|
|
|
|
return _AsyncCtx()
|
|
|
|
|
|
class TestFetchJWKS:
|
|
"""Coverage for ``fetch_jwks`` failure modes — malformed responses must
|
|
surface as ``OIDCError`` rather than ``AttributeError``/``KeyError``.
|
|
"""
|
|
|
|
def test_fetch_jwks_non_200_raises(self):
|
|
"""Non-2xx response (raise_for_status fires) -> OIDCError."""
|
|
from turnstone.core.oidc import fetch_jwks
|
|
|
|
mock_response = MagicMock()
|
|
mock_response.raise_for_status.side_effect = httpx.HTTPStatusError(
|
|
"404", request=MagicMock(), response=mock_response
|
|
)
|
|
|
|
async def _get(_url):
|
|
return mock_response
|
|
|
|
async def _run():
|
|
client = _mock_async_get_client(_get)
|
|
with (
|
|
patch("httpx.AsyncClient", return_value=client),
|
|
pytest.raises(OIDCError, match="JWKS fetch failed"),
|
|
):
|
|
await fetch_jwks("https://idp.example.com/jwks")
|
|
|
|
asyncio.run(_run())
|
|
|
|
def test_fetch_jwks_non_dict_body_raises(self):
|
|
"""Body decoded as a list rather than a JSON object -> OIDCError."""
|
|
from turnstone.core.oidc import fetch_jwks
|
|
|
|
mock_response = MagicMock()
|
|
mock_response.json.return_value = [1, 2, 3]
|
|
mock_response.raise_for_status = MagicMock()
|
|
|
|
async def _get(_url):
|
|
return mock_response
|
|
|
|
async def _run():
|
|
client = _mock_async_get_client(_get)
|
|
with (
|
|
patch("httpx.AsyncClient", return_value=client),
|
|
pytest.raises(OIDCError, match="not a JSON object"),
|
|
):
|
|
await fetch_jwks("https://idp.example.com/jwks")
|
|
|
|
asyncio.run(_run())
|
|
|
|
def test_fetch_jwks_dict_missing_keys_raises(self):
|
|
"""Body is a dict but lacks the ``keys`` array -> OIDCError."""
|
|
from turnstone.core.oidc import fetch_jwks
|
|
|
|
mock_response = MagicMock()
|
|
mock_response.json.return_value = {"not_keys": []}
|
|
mock_response.raise_for_status = MagicMock()
|
|
|
|
async def _get(_url):
|
|
return mock_response
|
|
|
|
async def _run():
|
|
client = _mock_async_get_client(_get)
|
|
with (
|
|
patch("httpx.AsyncClient", return_value=client),
|
|
pytest.raises(OIDCError, match="missing 'keys' array"),
|
|
):
|
|
await fetch_jwks("https://idp.example.com/jwks")
|
|
|
|
asyncio.run(_run())
|
|
|
|
def test_fetch_jwks_keys_not_a_list_raises(self):
|
|
"""``keys`` present but not a list -> OIDCError (not iterable type)."""
|
|
from turnstone.core.oidc import fetch_jwks
|
|
|
|
mock_response = MagicMock()
|
|
mock_response.json.return_value = {"keys": "definitely-not-a-list"}
|
|
mock_response.raise_for_status = MagicMock()
|
|
|
|
async def _get(_url):
|
|
return mock_response
|
|
|
|
async def _run():
|
|
client = _mock_async_get_client(_get)
|
|
with (
|
|
patch("httpx.AsyncClient", return_value=client),
|
|
pytest.raises(OIDCError, match="missing 'keys' array"),
|
|
):
|
|
await fetch_jwks("https://idp.example.com/jwks")
|
|
|
|
asyncio.run(_run())
|
|
|
|
def test_fetch_jwks_network_error_raises(self):
|
|
"""``httpx.RequestError`` from the transport -> OIDCError."""
|
|
from turnstone.core.oidc import fetch_jwks
|
|
|
|
async def _get(_url):
|
|
raise httpx.ConnectError("connection refused")
|
|
|
|
async def _run():
|
|
client = _mock_async_get_client(_get)
|
|
with (
|
|
patch("httpx.AsyncClient", return_value=client),
|
|
pytest.raises(OIDCError, match="JWKS fetch failed"),
|
|
):
|
|
await fetch_jwks("https://idp.example.com/jwks")
|
|
|
|
asyncio.run(_run())
|
|
|
|
|
|
class TestSanitizeLogText:
|
|
def test_sanitize_log_text_escapes_control_chars(self):
|
|
"""CR/LF, NUL, and tab characters must be escaped, not preserved."""
|
|
out = _sanitize_log_text("a\r\nb\tc\x00d", 500)
|
|
assert "\r" not in out
|
|
assert "\n" not in out
|
|
assert "\t" not in out
|
|
assert "\x00" not in out
|
|
assert "a" in out and "b" in out and "c" in out and "d" in out
|
|
|
|
def test_sanitize_log_text_truncates_after_escape(self):
|
|
"""The limit caps the rendered length, including the escape sequences."""
|
|
out = _sanitize_log_text("\n" * 100, 10)
|
|
assert len(out) == 10
|
|
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# User Provisioning
|
|
# ---------------------------------------------------------------------------
|
|
|
|
|
|
class TestProvisionOIDCUser:
|
|
def test_provision_oidc_user_existing(self):
|
|
"""Existing identity -> returns existing user, updates last_login."""
|
|
config = _make_config()
|
|
existing_user = {
|
|
"user_id": "u1",
|
|
"username": "alice",
|
|
"display_name": "Alice",
|
|
"password_hash": "!oidc",
|
|
}
|
|
existing_identity = {
|
|
"issuer": "https://idp.example.com",
|
|
"subject": "sub-123",
|
|
"user_id": "u1",
|
|
"email": "alice@example.com",
|
|
"created": "2024-01-01T00:00:00",
|
|
"last_login": "2024-01-01T00:00:00",
|
|
}
|
|
storage = _mock_storage(identity=existing_identity, user=existing_user)
|
|
|
|
claims = {"sub": "sub-123", "email": "alice@example.com", "name": "Alice"}
|
|
user = provision_oidc_user(storage, config, claims)
|
|
|
|
assert user["user_id"] == "u1"
|
|
assert user["username"] == "alice"
|
|
storage.update_oidc_identity_login.assert_called_once()
|
|
# Should not create a new user
|
|
storage.create_oidc_user.assert_not_called()
|
|
|
|
def test_provision_oidc_user_new(self):
|
|
"""No identity -> creates user + identity atomically."""
|
|
config = _make_config()
|
|
storage = _mock_storage()
|
|
|
|
# After create_oidc_user, get_user should return the new user
|
|
new_user = {
|
|
"user_id": "u-new",
|
|
"username": "bob",
|
|
"display_name": "Bob",
|
|
"password_hash": "!oidc",
|
|
}
|
|
storage.get_user.return_value = new_user
|
|
|
|
claims = {"sub": "sub-456", "preferred_username": "bob", "email": "bob@example.com"}
|
|
|
|
with patch("turnstone.core.oidc.uuid") as mock_uuid:
|
|
mock_uuid.uuid4.return_value = MagicMock(hex="u-new-hex-00000000000000000000")
|
|
user = provision_oidc_user(storage, config, claims)
|
|
|
|
assert user["username"] == "bob"
|
|
storage.create_oidc_user.assert_called_once()
|
|
# Positional args: user_id, username, display_name, password_hash, issuer, subject, email
|
|
call_args = storage.create_oidc_user.call_args
|
|
assert call_args[0][1] == "bob"
|
|
assert call_args[0][3] == "!oidc"
|
|
assert call_args[0][4] == "https://idp.example.com"
|
|
assert call_args[0][5] == "sub-456"
|
|
assert call_args[0][6] == "bob@example.com"
|
|
|
|
def test_provision_oidc_user_username_dedup(self):
|
|
"""First username taken -> appends suffix."""
|
|
config = _make_config()
|
|
# Bulk lookup reports "bob" already taken; "bob2" is free.
|
|
storage = _mock_storage(existing_usernames={"bob"})
|
|
new_user = {
|
|
"user_id": "u-new",
|
|
"username": "bob2",
|
|
"display_name": "Bob",
|
|
"password_hash": "!oidc",
|
|
}
|
|
storage.get_user.return_value = new_user
|
|
|
|
claims = {"sub": "sub-789", "preferred_username": "bob", "email": "bob@example.com"}
|
|
user = provision_oidc_user(storage, config, claims)
|
|
|
|
assert user["username"] == "bob2"
|
|
# create_oidc_user should have been called with "bob2" as username
|
|
call_args = storage.create_oidc_user.call_args
|
|
assert call_args[0][1] == "bob2"
|
|
# Single bulk query rather than per-candidate get_user_by_username
|
|
storage.find_existing_usernames.assert_called_once()
|
|
storage.get_user_by_username.assert_not_called()
|
|
|
|
def test_provision_oidc_user_email_prefix(self):
|
|
"""No preferred_username -> uses email prefix."""
|
|
config = _make_config()
|
|
storage = _mock_storage()
|
|
|
|
new_user = {
|
|
"user_id": "u-new",
|
|
"username": "charlie",
|
|
"display_name": "charlie@example.com",
|
|
"password_hash": "!oidc",
|
|
}
|
|
storage.get_user.return_value = new_user
|
|
|
|
claims = {"sub": "sub-abc", "email": "charlie@example.com"}
|
|
provision_oidc_user(storage, config, claims)
|
|
|
|
call_args = storage.create_oidc_user.call_args
|
|
assert call_args[0][1] == "charlie"
|
|
|
|
def test_provision_oidc_user_missing_user_raises(self):
|
|
"""Identity references missing user -> raises OIDCError."""
|
|
config = _make_config()
|
|
existing_identity = {
|
|
"issuer": "https://idp.example.com",
|
|
"subject": "sub-orphan",
|
|
"user_id": "u-gone",
|
|
"email": "gone@example.com",
|
|
"created": "2024-01-01T00:00:00",
|
|
"last_login": "2024-01-01T00:00:00",
|
|
}
|
|
storage = _mock_storage(identity=existing_identity, user=None)
|
|
|
|
claims = {"sub": "sub-orphan", "email": "gone@example.com"}
|
|
with pytest.raises(OIDCError, match="missing user"):
|
|
provision_oidc_user(storage, config, claims)
|
|
|
|
def test_provision_oidc_user_fallback_username(self):
|
|
"""No preferred_username and no email -> falls back to 'user'."""
|
|
config = _make_config()
|
|
storage = _mock_storage()
|
|
new_user = {
|
|
"user_id": "u-new",
|
|
"username": "user",
|
|
"display_name": "",
|
|
"password_hash": "!oidc",
|
|
}
|
|
storage.get_user.return_value = new_user
|
|
|
|
claims = {"sub": "sub-noemail"}
|
|
provision_oidc_user(storage, config, claims)
|
|
|
|
call_args = storage.create_oidc_user.call_args
|
|
assert call_args[0][1] == "user"
|
|
|
|
def test_provision_oidc_user_username_conflict_raises(self):
|
|
"""Username TOCTOU -> create_oidc_user raises StorageConflictError;
|
|
provision_oidc_user wraps it as OIDCError without leaving role rows.
|
|
"""
|
|
from turnstone.core.storage import StorageConflictError
|
|
|
|
config = _make_config()
|
|
storage = _mock_storage()
|
|
storage.create_oidc_user.side_effect = StorageConflictError("username already taken: bob")
|
|
|
|
claims = {"sub": "sub-race", "preferred_username": "bob", "email": "bob@example.com"}
|
|
with pytest.raises(OIDCError, match="username already taken"):
|
|
provision_oidc_user(storage, config, claims)
|
|
|
|
storage.assign_role.assert_not_called()
|
|
|
|
def test_provision_oidc_user_identity_conflict_raises(self):
|
|
"""Concurrent (issuer, subject) creates -> StorageConflictError -> OIDCError."""
|
|
from turnstone.core.storage import StorageConflictError
|
|
|
|
config = _make_config()
|
|
storage = _mock_storage()
|
|
storage.create_oidc_user.side_effect = StorageConflictError(
|
|
"OIDC identity already linked: (https://idp.example.com, sub-race)"
|
|
)
|
|
|
|
claims = {"sub": "sub-race", "preferred_username": "bob", "email": "bob@example.com"}
|
|
with pytest.raises(OIDCError, match="OIDC identity already linked"):
|
|
provision_oidc_user(storage, config, claims)
|
|
|
|
storage.assign_role.assert_not_called()
|
|
|
|
def test_provision_oidc_user_null_oid_tid_collapse_to_empty(self):
|
|
"""A present-but-null oid/tid claim must store "" — never the string "None".
|
|
|
|
`claims.get("oid", "")` returns None (not the "" default) when the key is
|
|
present with a JSON null, and str(None) == "None" would slip past both the
|
|
server_default and the truthy backfill guard, storing a bogus non-empty
|
|
sentinel that collides across every null-emitting user. New-user path.
|
|
"""
|
|
config = _make_config()
|
|
storage = _mock_storage()
|
|
storage.get_user.return_value = {
|
|
"user_id": "u-new",
|
|
"username": "bob",
|
|
"display_name": "Bob",
|
|
"password_hash": "!oidc",
|
|
}
|
|
|
|
claims = {"sub": "sub-null", "preferred_username": "bob", "oid": None, "tid": None}
|
|
with patch("turnstone.core.oidc.uuid") as mock_uuid:
|
|
mock_uuid.uuid4.return_value = MagicMock(hex="u-new-hex-00000000000000000000")
|
|
provision_oidc_user(storage, config, claims)
|
|
|
|
kwargs = storage.create_oidc_user.call_args.kwargs
|
|
assert kwargs["oid"] == ""
|
|
assert kwargs["tid"] == ""
|
|
|
|
def test_provision_oidc_user_null_oid_tid_not_backfilled_existing(self):
|
|
"""Existing-identity path: null oid/tid claims must not backfill "None".
|
|
|
|
The truthy guard in update_oidc_identity_login only protects against ""; a
|
|
"None" produced by str(None) is truthy and would be written, clobbering a
|
|
real value captured on an earlier login.
|
|
"""
|
|
config = _make_config()
|
|
existing_user = {
|
|
"user_id": "u1",
|
|
"username": "alice",
|
|
"display_name": "Alice",
|
|
"password_hash": "!oidc",
|
|
}
|
|
existing_identity = {
|
|
"issuer": "https://idp.example.com",
|
|
"subject": "sub-123",
|
|
"user_id": "u1",
|
|
"email": "alice@example.com",
|
|
"created": "2024-01-01T00:00:00",
|
|
"last_login": "2024-01-01T00:00:00",
|
|
"oid": "obj-real",
|
|
"tid": "ten-real",
|
|
}
|
|
storage = _mock_storage(identity=existing_identity, user=existing_user)
|
|
|
|
claims = {"sub": "sub-123", "email": "alice@example.com", "oid": None, "tid": None}
|
|
provision_oidc_user(storage, config, claims)
|
|
|
|
kwargs = storage.update_oidc_identity_login.call_args.kwargs
|
|
assert kwargs["oid"] == ""
|
|
assert kwargs["tid"] == ""
|
|
|
|
def test_existing_identity_self_heals_zero_roles(self):
|
|
"""Existing identity user with zero roles -> safety-net assigns builtin-viewer.
|
|
|
|
Models the bug-1 strand: a prior login committed user + identity but
|
|
``apply_role_mapping`` raised before reaching the safety-net. On the
|
|
next login the existing-identity branch must self-heal.
|
|
"""
|
|
config = _make_config()
|
|
existing_user = {
|
|
"user_id": "u-stranded",
|
|
"username": "alice",
|
|
"display_name": "Alice",
|
|
"password_hash": "!oidc",
|
|
}
|
|
existing_identity = {
|
|
"issuer": "https://idp.example.com",
|
|
"subject": "sub-stranded",
|
|
"user_id": "u-stranded",
|
|
"email": "alice@example.com",
|
|
"created": "2024-01-01T00:00:00",
|
|
"last_login": "2024-01-01T00:00:00",
|
|
}
|
|
storage = _mock_storage(
|
|
identity=existing_identity,
|
|
user=existing_user,
|
|
role={"role_id": "builtin-viewer", "name": "Viewer"},
|
|
)
|
|
storage.list_user_roles.return_value = []
|
|
|
|
claims = {"sub": "sub-stranded", "email": "alice@example.com", "name": "Alice"}
|
|
user = provision_oidc_user(storage, config, claims)
|
|
|
|
assert user["user_id"] == "u-stranded"
|
|
storage.list_user_roles.assert_called_once_with("u-stranded")
|
|
storage.assign_role.assert_called_once_with("u-stranded", "builtin-viewer", "oidc-default")
|
|
|
|
def test_existing_identity_does_not_re_assign_when_user_has_roles(self):
|
|
"""User already has at least one role -> safety-net no-ops."""
|
|
config = _make_config()
|
|
existing_user = {
|
|
"user_id": "u1",
|
|
"username": "alice",
|
|
"display_name": "Alice",
|
|
"password_hash": "!oidc",
|
|
}
|
|
existing_identity = {
|
|
"issuer": "https://idp.example.com",
|
|
"subject": "sub-123",
|
|
"user_id": "u1",
|
|
"email": "alice@example.com",
|
|
"created": "2024-01-01T00:00:00",
|
|
"last_login": "2024-01-01T00:00:00",
|
|
}
|
|
storage = _mock_storage(
|
|
identity=existing_identity,
|
|
user=existing_user,
|
|
role={"role_id": "builtin-viewer", "name": "Viewer"},
|
|
)
|
|
storage.list_user_roles.return_value = [{"role_id": "builtin-operator"}]
|
|
|
|
claims = {"sub": "sub-123", "email": "alice@example.com"}
|
|
provision_oidc_user(storage, config, claims)
|
|
|
|
storage.list_user_roles.assert_called_once_with("u1")
|
|
storage.assign_role.assert_not_called()
|
|
|
|
def test_existing_identity_with_claim_mapped_roles_skips_default(self):
|
|
"""Claim-driven mapping populated roles -> hint short-circuits the helper."""
|
|
config = _make_config(
|
|
role_claim="groups",
|
|
role_map={"admin": "builtin-admin"},
|
|
)
|
|
existing_user = {
|
|
"user_id": "u1",
|
|
"username": "alice",
|
|
"display_name": "Alice",
|
|
"password_hash": "!oidc",
|
|
}
|
|
existing_identity = {
|
|
"issuer": "https://idp.example.com",
|
|
"subject": "sub-123",
|
|
"user_id": "u1",
|
|
"email": "alice@example.com",
|
|
"created": "2024-01-01T00:00:00",
|
|
"last_login": "2024-01-01T00:00:00",
|
|
}
|
|
storage = _mock_storage(
|
|
identity=existing_identity,
|
|
user=existing_user,
|
|
role={"role_id": "builtin-admin", "name": "Admin"},
|
|
)
|
|
|
|
claims = {"sub": "sub-123", "groups": "admin"}
|
|
provision_oidc_user(storage, config, claims)
|
|
|
|
storage.list_user_roles.assert_not_called()
|
|
storage.assign_role.assert_not_called()
|
|
|
|
def test_new_user_safety_net_still_fires(self):
|
|
"""Fresh user with no IdP-mapped roles -> builtin-viewer fallback applied."""
|
|
config = _make_config()
|
|
storage = _mock_storage(role={"role_id": "builtin-viewer", "name": "Viewer"})
|
|
storage.list_user_roles.return_value = []
|
|
new_user = {
|
|
"user_id": "u-new",
|
|
"username": "bob",
|
|
"display_name": "Bob",
|
|
"password_hash": "!oidc",
|
|
}
|
|
storage.get_user.return_value = new_user
|
|
|
|
claims = {"sub": "sub-new", "preferred_username": "bob", "email": "bob@example.com"}
|
|
provision_oidc_user(storage, config, claims)
|
|
|
|
storage.assign_role.assert_called_once()
|
|
called_args = storage.assign_role.call_args[0]
|
|
assert called_args[1] == "builtin-viewer"
|
|
assert called_args[2] == "oidc-default"
|
|
|
|
def test_new_user_safety_net_skipped_when_apply_role_mapping_assigns_role(self):
|
|
"""Claim-driven mapping populates roles -> safety-net hint short-circuits."""
|
|
config = _make_config(
|
|
role_claim="groups",
|
|
role_map={"admin": "builtin-admin"},
|
|
)
|
|
storage = _mock_storage(role={"role_id": "builtin-admin", "name": "Admin"})
|
|
new_user = {
|
|
"user_id": "u-new",
|
|
"username": "bob",
|
|
"display_name": "Bob",
|
|
"password_hash": "!oidc",
|
|
}
|
|
storage.get_user.return_value = new_user
|
|
|
|
claims = {
|
|
"sub": "sub-new",
|
|
"preferred_username": "bob",
|
|
"email": "bob@example.com",
|
|
"groups": "admin",
|
|
}
|
|
provision_oidc_user(storage, config, claims)
|
|
|
|
storage.list_user_roles.assert_not_called()
|
|
storage.assign_role.assert_not_called()
|
|
|
|
def test_ensure_default_role_noop_when_builtin_viewer_missing(self):
|
|
"""builtin-viewer absent from role table -> helper does nothing."""
|
|
storage = _mock_storage(role=None)
|
|
|
|
_ensure_default_role(storage, "u1")
|
|
|
|
storage.get_role.assert_called_once_with("builtin-viewer")
|
|
storage.list_user_roles.assert_not_called()
|
|
storage.assign_role.assert_not_called()
|
|
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# Username derivation — UUID-retry fallback tiers
|
|
# ---------------------------------------------------------------------------
|
|
|
|
|
|
class TestDeriveUsername:
|
|
"""Tier-3 (UUID-retry) and tier-4 (give-up) coverage for ``_derive_username``.
|
|
|
|
Tier 1 (sanitised candidate) and tier 2 (numeric suffix dedup) are
|
|
already exercised through ``TestProvisionOIDCUser``; these tests
|
|
drive the post-batch-6 fallback path that runs after every numeric
|
|
suffix collides.
|
|
"""
|
|
|
|
def _claims(self) -> dict[str, str]:
|
|
return {
|
|
"sub": "sub-x",
|
|
"preferred_username": "bob",
|
|
"email": "bob@example.com",
|
|
}
|
|
|
|
def _all_numeric_candidates_taken(self) -> set[str]:
|
|
"""The full set of tier-1 + tier-2 candidates ``_derive_username`` enumerates."""
|
|
return {"bob", *(f"bob{n}" for n in range(2, 11))}
|
|
|
|
def test_derive_username_falls_into_uuid_retry_when_all_suffixes_taken(self):
|
|
"""All 10 numeric candidates taken -> tier 3 returns a UUID-suffixed username."""
|
|
config = _make_config()
|
|
storage = _mock_storage(existing_usernames=self._all_numeric_candidates_taken())
|
|
# First UUID candidate is free.
|
|
storage.get_user_by_username.return_value = None
|
|
new_user = {
|
|
"user_id": "u-new",
|
|
"username": "ignored-by-test",
|
|
"display_name": "Bob",
|
|
"password_hash": "!oidc",
|
|
}
|
|
storage.get_user.return_value = new_user
|
|
|
|
provision_oidc_user(storage, config, self._claims())
|
|
|
|
# The username actually persisted is the tier-3 UUID candidate.
|
|
called_with = storage.create_oidc_user.call_args[0][1]
|
|
assert called_with.startswith("bob")
|
|
# bob + 32-hex UUID hex == 35 chars total.
|
|
suffix = called_with[len("bob") :]
|
|
assert len(suffix) == 32
|
|
assert all(c in "0123456789abcdef" for c in suffix)
|
|
# Exactly one tier-3 lookup happened (no retry collision).
|
|
assert storage.get_user_by_username.call_count == 1
|
|
|
|
def test_derive_username_uuid_retry_succeeds_on_second_attempt(self):
|
|
"""First UUID candidate collides; second attempt is free."""
|
|
config = _make_config()
|
|
storage = _mock_storage(existing_usernames=self._all_numeric_candidates_taken())
|
|
# First lookup: collision (existing user); second: free.
|
|
storage.get_user_by_username.side_effect = [
|
|
{"user_id": "other", "username": "bob-already-used"},
|
|
None,
|
|
]
|
|
new_user = {
|
|
"user_id": "u-new",
|
|
"username": "ignored-by-test",
|
|
"display_name": "Bob",
|
|
"password_hash": "!oidc",
|
|
}
|
|
storage.get_user.return_value = new_user
|
|
|
|
provision_oidc_user(storage, config, self._claims())
|
|
|
|
assert storage.get_user_by_username.call_count == 2
|
|
# The username persisted is the second UUID candidate.
|
|
persisted = storage.create_oidc_user.call_args[0][1]
|
|
assert persisted.startswith("bob")
|
|
assert len(persisted) == len("bob") + 32
|
|
|
|
def test_derive_username_uuid_retry_exhausted_raises(self):
|
|
"""All 3 UUID-retry attempts collide -> OIDCError raised."""
|
|
config = _make_config()
|
|
storage = _mock_storage(existing_usernames=self._all_numeric_candidates_taken())
|
|
# Every UUID candidate hits an existing user.
|
|
storage.get_user_by_username.return_value = {"user_id": "other", "username": "x"}
|
|
|
|
with pytest.raises(OIDCError, match="Failed to generate unique username"):
|
|
provision_oidc_user(storage, config, self._claims())
|
|
|
|
# Tier 3 is bounded to 3 attempts.
|
|
assert storage.get_user_by_username.call_count == 3
|
|
storage.create_oidc_user.assert_not_called()
|
|
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# Role Mapping
|
|
# ---------------------------------------------------------------------------
|
|
|
|
|
|
class TestApplyRoleMapping:
|
|
def test_apply_role_mapping_basic(self):
|
|
"""Maps claim value to role."""
|
|
config = _make_config(
|
|
role_claim="groups",
|
|
role_map={"admin": "builtin-admin"},
|
|
)
|
|
storage = _mock_storage(role={"role_id": "builtin-admin", "name": "Admin"})
|
|
|
|
claims = {"sub": "u1", "groups": "admin"}
|
|
apply_role_mapping(storage, "u1", claims, config)
|
|
|
|
storage.replace_oidc_roles.assert_called_once_with("u1", {"builtin-admin"})
|
|
storage.assign_role.assert_not_called()
|
|
storage.unassign_role.assert_not_called()
|
|
storage.list_user_roles.assert_not_called()
|
|
|
|
def test_apply_role_mapping_list_claim(self):
|
|
"""Claim is a list of strings -> single replace call."""
|
|
config = _make_config(
|
|
role_claim="roles",
|
|
role_map={"admin": "builtin-admin", "editor": "builtin-operator"},
|
|
)
|
|
storage = _mock_storage()
|
|
# get_role returns non-None for both roles
|
|
storage.get_role.return_value = {"role_id": "some-role"}
|
|
|
|
claims = {"sub": "u1", "roles": ["admin", "editor"]}
|
|
apply_role_mapping(storage, "u1", claims, config)
|
|
|
|
storage.replace_oidc_roles.assert_called_once_with(
|
|
"u1", {"builtin-admin", "builtin-operator"}
|
|
)
|
|
|
|
def test_apply_role_mapping_no_config(self):
|
|
"""No role_claim configured -> no-op."""
|
|
config = _make_config(role_claim="", role_map={})
|
|
storage = _mock_storage()
|
|
|
|
claims = {"sub": "u1", "roles": "admin"}
|
|
apply_role_mapping(storage, "u1", claims, config)
|
|
|
|
storage.replace_oidc_roles.assert_not_called()
|
|
|
|
def test_apply_role_mapping_unknown_role(self):
|
|
"""Claim maps to nonexistent role -> empty replace set."""
|
|
config = _make_config(
|
|
role_claim="groups",
|
|
role_map={"admin": "nonexistent-role"},
|
|
)
|
|
storage = _mock_storage(role=None) # role doesn't exist
|
|
|
|
claims = {"sub": "u1", "groups": "admin"}
|
|
apply_role_mapping(storage, "u1", claims, config)
|
|
|
|
storage.replace_oidc_roles.assert_called_once_with("u1", set())
|
|
|
|
def test_apply_role_mapping_no_matching_claim_value(self):
|
|
"""Claim value not in role_map -> empty replace set."""
|
|
config = _make_config(
|
|
role_claim="groups",
|
|
role_map={"admin": "builtin-admin"},
|
|
)
|
|
storage = _mock_storage()
|
|
|
|
claims = {"sub": "u1", "groups": "viewer"} # "viewer" not in role_map
|
|
apply_role_mapping(storage, "u1", claims, config)
|
|
|
|
storage.replace_oidc_roles.assert_called_once_with("u1", set())
|
|
|
|
def test_apply_role_mapping_claim_missing(self):
|
|
"""Claim key not present in claims -> empty replace set."""
|
|
config = _make_config(
|
|
role_claim="groups",
|
|
role_map={"admin": "builtin-admin"},
|
|
)
|
|
storage = _mock_storage()
|
|
|
|
claims = {"sub": "u1"} # no "groups" key
|
|
apply_role_mapping(storage, "u1", claims, config)
|
|
|
|
storage.replace_oidc_roles.assert_called_once_with("u1", set())
|
|
|
|
def test_apply_role_mapping_no_role_map(self):
|
|
"""role_claim set but role_map empty -> no-op (early return)."""
|
|
config = _make_config(role_claim="groups", role_map={})
|
|
storage = _mock_storage()
|
|
|
|
claims = {"sub": "u1", "groups": "admin"}
|
|
apply_role_mapping(storage, "u1", claims, config)
|
|
|
|
storage.replace_oidc_roles.assert_not_called()
|
|
|
|
def test_apply_role_mapping_revokes_stale_oidc_roles(self, caplog):
|
|
"""Roles previously assigned by OIDC but no longer in claims are revoked.
|
|
|
|
Storage owns the diff via ``replace_oidc_roles``; ``apply_role_mapping``
|
|
only logs the returned ``(added, removed)`` sets. Manual roles are
|
|
invisible to ``apply_role_mapping`` post-perf-3 — protection now lives
|
|
in the storage layer's ``WHERE assigned_by = 'oidc'`` filter.
|
|
"""
|
|
config = _make_config(
|
|
role_claim="groups",
|
|
role_map={"admin": "builtin-admin", "eng": "builtin-operator"},
|
|
)
|
|
storage = _mock_storage(
|
|
replace_oidc_roles=({"builtin-operator"}, {"builtin-admin"}),
|
|
)
|
|
storage.get_role.return_value = {"role_id": "some-role"}
|
|
|
|
# IdP now only says "eng", not "admin"
|
|
claims = {"sub": "u1", "groups": ["eng"]}
|
|
with caplog.at_level("INFO", logger="turnstone.core.oidc"):
|
|
apply_role_mapping(storage, "u1", claims, config)
|
|
|
|
storage.replace_oidc_roles.assert_called_once_with("u1", {"builtin-operator"})
|
|
assert any(
|
|
"Revoked role" in record.getMessage() and "builtin-admin" in record.getMessage()
|
|
for record in caplog.records
|
|
)
|
|
|
|
def test_apply_role_mapping_revokes_all_oidc_roles_when_claim_absent(self):
|
|
"""Empty desired set propagates to storage (diff happens server-side)."""
|
|
config = _make_config(
|
|
role_claim="groups",
|
|
role_map={"admin": "builtin-admin"},
|
|
)
|
|
storage = _mock_storage(
|
|
replace_oidc_roles=(set(), {"builtin-admin", "builtin-operator"}),
|
|
)
|
|
storage.get_role.return_value = {"role_id": "some-role"}
|
|
|
|
claims = {"sub": "u1"} # no "groups" key
|
|
apply_role_mapping(storage, "u1", claims, config)
|
|
|
|
storage.replace_oidc_roles.assert_called_once_with("u1", set())
|
|
|
|
def test_apply_role_mapping_int_claim(self):
|
|
"""Numeric claim values stringify before role_map lookup (no crash)."""
|
|
config = _make_config(
|
|
role_claim="roles",
|
|
role_map={"42": "builtin-viewer"},
|
|
)
|
|
storage = _mock_storage(role={"role_id": "builtin-viewer", "name": "Viewer"})
|
|
|
|
claims = {"sub": "u1", "roles": 42}
|
|
apply_role_mapping(storage, "u1", claims, config)
|
|
|
|
storage.replace_oidc_roles.assert_called_once_with("u1", {"builtin-viewer"})
|
|
|
|
def test_apply_role_mapping_dict_claim(self):
|
|
"""Dict claim values stringify but won't reliably hit role_map -> empty set."""
|
|
config = _make_config(
|
|
role_claim="roles",
|
|
role_map={"admin": "builtin-admin"},
|
|
)
|
|
storage = _mock_storage()
|
|
|
|
# Dict reprs are unstable across runtime versions, so callers cannot
|
|
# reasonably configure role_map keys to match. The contract is just
|
|
# "don't crash".
|
|
claims = {"sub": "u1", "roles": {"role": "admin"}}
|
|
apply_role_mapping(storage, "u1", claims, config)
|
|
|
|
storage.replace_oidc_roles.assert_called_once_with("u1", set())
|
|
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# Discovery (async)
|
|
# ---------------------------------------------------------------------------
|
|
|
|
|
|
class TestDiscoverOIDC:
|
|
# Mock DNS result for a public IP — reused across discovery tests.
|
|
_PUBLIC_ADDR = [(2, 1, 6, "", ("93.184.216.34", 0))]
|
|
|
|
def test_discover_oidc_success(self):
|
|
"""Mock httpx response, verify endpoints populated."""
|
|
config = _make_config(
|
|
authorization_endpoint="",
|
|
token_endpoint="",
|
|
userinfo_endpoint="",
|
|
jwks_uri="",
|
|
)
|
|
|
|
discovery_doc = {
|
|
"authorization_endpoint": "https://idp.example.com/authorize",
|
|
"token_endpoint": "https://idp.example.com/token",
|
|
"userinfo_endpoint": "https://idp.example.com/userinfo",
|
|
"jwks_uri": "https://idp.example.com/.well-known/jwks.json",
|
|
}
|
|
|
|
mock_response = MagicMock()
|
|
mock_response.json.return_value = discovery_doc
|
|
mock_response.raise_for_status = MagicMock()
|
|
|
|
async def _run():
|
|
client = _mock_async_client(lambda url: _async_return(mock_response))
|
|
with (
|
|
patch("socket.getaddrinfo", return_value=self._PUBLIC_ADDR),
|
|
patch("httpx.AsyncClient", return_value=client),
|
|
):
|
|
result = await discover_oidc(config)
|
|
|
|
assert result.authorization_endpoint == "https://idp.example.com/authorize"
|
|
assert result.token_endpoint == "https://idp.example.com/token"
|
|
assert result.userinfo_endpoint == "https://idp.example.com/userinfo"
|
|
assert result.jwks_uri == "https://idp.example.com/.well-known/jwks.json"
|
|
assert result.enabled is True
|
|
|
|
asyncio.run(_run())
|
|
|
|
def test_discover_oidc_private_issuer_rejected_by_default(self):
|
|
"""Without the opt-in, a private-resolving issuer disables OIDC."""
|
|
config = _make_config(
|
|
issuer="https://auth.internal.example",
|
|
authorization_endpoint="",
|
|
token_endpoint="",
|
|
userinfo_endpoint="",
|
|
jwks_uri="",
|
|
)
|
|
|
|
async def _run():
|
|
with patch("socket.getaddrinfo", return_value=[(2, 1, 6, "", ("10.0.0.5", 0))]):
|
|
result = await discover_oidc(config)
|
|
assert result.enabled is False
|
|
|
|
asyncio.run(_run())
|
|
|
|
def test_discover_oidc_private_issuer_with_opt_in(self):
|
|
"""allow_private_network=True lets a private-resolving IdP discover."""
|
|
config = _make_config(
|
|
issuer="https://auth.internal.example",
|
|
allow_private_network=True,
|
|
authorization_endpoint="",
|
|
token_endpoint="",
|
|
userinfo_endpoint="",
|
|
jwks_uri="",
|
|
)
|
|
|
|
discovery_doc = {
|
|
"authorization_endpoint": "https://auth.internal.example/authorize",
|
|
"token_endpoint": "https://auth.internal.example/token",
|
|
"userinfo_endpoint": "https://auth.internal.example/userinfo",
|
|
"jwks_uri": "https://auth.internal.example/jwks",
|
|
}
|
|
|
|
mock_response = MagicMock()
|
|
mock_response.json.return_value = discovery_doc
|
|
mock_response.raise_for_status = MagicMock()
|
|
|
|
async def _run():
|
|
client = _mock_async_client(lambda url: _async_return(mock_response))
|
|
with (
|
|
patch("socket.getaddrinfo", return_value=[(2, 1, 6, "", ("10.0.0.5", 0))]),
|
|
patch("httpx.AsyncClient", return_value=client),
|
|
):
|
|
result = await discover_oidc(config)
|
|
|
|
assert result.enabled is True
|
|
assert result.token_endpoint == "https://auth.internal.example/token"
|
|
|
|
asyncio.run(_run())
|
|
|
|
def test_discover_oidc_failure(self):
|
|
"""Mock httpx error -> enabled=False returned."""
|
|
config = _make_config(
|
|
authorization_endpoint="",
|
|
token_endpoint="",
|
|
userinfo_endpoint="",
|
|
jwks_uri="",
|
|
)
|
|
|
|
async def _failing_get(url):
|
|
raise httpx.ConnectError("connection refused")
|
|
|
|
async def _run():
|
|
client = _mock_async_client(_failing_get)
|
|
with (
|
|
patch("socket.getaddrinfo", return_value=self._PUBLIC_ADDR),
|
|
patch("httpx.AsyncClient", return_value=client),
|
|
):
|
|
result = await discover_oidc(config)
|
|
|
|
assert result.enabled is False
|
|
|
|
asyncio.run(_run())
|
|
|
|
def test_discover_oidc_no_issuer(self):
|
|
"""Empty issuer -> enabled=False."""
|
|
config = _make_config(issuer="")
|
|
|
|
async def _run():
|
|
result = await discover_oidc(config)
|
|
assert result.enabled is False
|
|
|
|
asyncio.run(_run())
|
|
|
|
def test_discover_oidc_missing_required_endpoints(self):
|
|
"""Discovery doc missing authorization_endpoint -> enabled=False."""
|
|
config = _make_config(
|
|
authorization_endpoint="",
|
|
token_endpoint="",
|
|
userinfo_endpoint="",
|
|
jwks_uri="",
|
|
)
|
|
|
|
# Document missing authorization_endpoint
|
|
discovery_doc = {
|
|
"token_endpoint": "https://idp.example.com/token",
|
|
"jwks_uri": "https://idp.example.com/.well-known/jwks.json",
|
|
}
|
|
|
|
mock_response = MagicMock()
|
|
mock_response.json.return_value = discovery_doc
|
|
mock_response.raise_for_status = MagicMock()
|
|
|
|
async def _run():
|
|
client = _mock_async_client(lambda url: _async_return(mock_response))
|
|
with (
|
|
patch("socket.getaddrinfo", return_value=self._PUBLIC_ADDR),
|
|
patch("httpx.AsyncClient", return_value=client),
|
|
):
|
|
result = await discover_oidc(config)
|
|
|
|
assert result.enabled is False
|
|
|
|
asyncio.run(_run())
|
|
|
|
def test_discover_oidc_handles_non_dict_response(self):
|
|
"""IdP returning a list/null body must not raise AttributeError."""
|
|
config = _make_config(
|
|
authorization_endpoint="",
|
|
token_endpoint="",
|
|
userinfo_endpoint="",
|
|
jwks_uri="",
|
|
)
|
|
|
|
for payload in (["not", "a", "dict"], None, "string-body", 42):
|
|
mock_response = MagicMock()
|
|
mock_response.json.return_value = payload
|
|
mock_response.raise_for_status = MagicMock()
|
|
|
|
async def _run(resp=mock_response):
|
|
client = _mock_async_client(lambda url: _async_return(resp))
|
|
with (
|
|
patch("socket.getaddrinfo", return_value=self._PUBLIC_ADDR),
|
|
patch("httpx.AsyncClient", return_value=client),
|
|
):
|
|
return await discover_oidc(config)
|
|
|
|
result = asyncio.run(_run())
|
|
assert result.enabled is False
|
|
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# Runtime re-discovery (boot-time transient failure self-heal)
|
|
# ---------------------------------------------------------------------------
|
|
|
|
|
|
class TestRuntimeRediscovery:
|
|
"""A node that boots during a transient IdP outage keeps enabled=False
|
|
forever without a runtime retry — OIDC login stays dark and every
|
|
oauth_obo mint on the node fails "transient" until an operator restarts
|
|
it. ``maybe_rediscover_oidc`` heals that, cooldown-gated; config-caused
|
|
discovery failures (bad issuer, SSRF rejection) are NOT retried."""
|
|
|
|
_PUBLIC_ADDR = [(2, 1, 6, "", ("93.184.216.34", 0))]
|
|
|
|
def test_fetch_failure_marks_config_retryable(self):
|
|
"""The transient branch (IdP unreachable) sets discovery_retryable."""
|
|
config = _make_config(
|
|
authorization_endpoint="", token_endpoint="", userinfo_endpoint="", jwks_uri=""
|
|
)
|
|
|
|
async def _raise_connect(url):
|
|
raise httpx.ConnectError("boom")
|
|
|
|
async def _run():
|
|
client = _mock_async_client(_raise_connect)
|
|
with (
|
|
patch("socket.getaddrinfo", return_value=self._PUBLIC_ADDR),
|
|
patch("httpx.AsyncClient", return_value=client),
|
|
):
|
|
return await discover_oidc(config)
|
|
|
|
result = asyncio.run(_run())
|
|
assert result.enabled is False
|
|
assert result.discovery_retryable is True
|
|
|
|
def test_config_rejection_is_not_retryable(self):
|
|
"""An SSRF-rejected issuer is a config problem — retrying is pointless."""
|
|
config = _make_config(
|
|
issuer="https://auth.internal.example",
|
|
authorization_endpoint="",
|
|
token_endpoint="",
|
|
userinfo_endpoint="",
|
|
jwks_uri="",
|
|
)
|
|
|
|
async def _run():
|
|
with patch("socket.getaddrinfo", return_value=[(2, 1, 6, "", ("10.0.0.5", 0))]):
|
|
return await discover_oidc(config)
|
|
|
|
result = asyncio.run(_run())
|
|
assert result.enabled is False
|
|
assert result.discovery_retryable is False
|
|
|
|
def _disabled_retryable_state(self) -> SimpleNamespace:
|
|
cfg = _make_config(
|
|
enabled=False,
|
|
authorization_endpoint="",
|
|
token_endpoint="",
|
|
userinfo_endpoint="",
|
|
jwks_uri="",
|
|
)
|
|
cfg = dataclasses.replace(cfg, discovery_retryable=True)
|
|
return SimpleNamespace(oidc_config=cfg)
|
|
|
|
def test_rediscover_swaps_enabled_config_on_success(self):
|
|
"""Drives the REAL discover_oidc through a mocked HTTP discovery GET
|
|
(NOT a mock of discover_oidc itself): discover_oidc preserves the input
|
|
config's ``enabled`` on success and only clears it on failure, so a
|
|
probe started from the disabled boot config must first force enabled=True
|
|
or the recovered config never installs. An earlier version of this test
|
|
mocked discover_oidc to return enabled=True and so masked exactly that
|
|
dead-code bug."""
|
|
from turnstone.core.oidc import maybe_rediscover_oidc
|
|
|
|
state = self._disabled_retryable_state() # issuer=https://idp.example.com
|
|
discovery_doc = {
|
|
"authorization_endpoint": "https://idp.example.com/authorize",
|
|
"token_endpoint": "https://idp.example.com/token",
|
|
"userinfo_endpoint": "https://idp.example.com/userinfo",
|
|
"jwks_uri": "https://idp.example.com/.well-known/jwks.json",
|
|
}
|
|
mock_response = MagicMock()
|
|
mock_response.json.return_value = discovery_doc
|
|
mock_response.raise_for_status = MagicMock()
|
|
|
|
async def _run():
|
|
client = _mock_async_client(lambda url: _async_return(mock_response))
|
|
with (
|
|
patch("socket.getaddrinfo", return_value=[(2, 1, 6, "", ("93.184.216.34", 0))]),
|
|
patch("httpx.AsyncClient", return_value=client),
|
|
):
|
|
await maybe_rediscover_oidc(state)
|
|
|
|
asyncio.run(_run())
|
|
|
|
# The real discover_oidc succeeded and the recovered config was installed.
|
|
assert state.oidc_config.enabled is True
|
|
assert state.oidc_config.token_endpoint == "https://idp.example.com/token"
|
|
# The healed config no longer advertises a retryable failure.
|
|
assert state.oidc_config.discovery_retryable is False
|
|
|
|
def test_rediscover_latches_terminal_on_config_error_and_stops_probing(self):
|
|
"""Review finding: probing with enabled forced True carries the
|
|
retryable boot flag into discover_oidc, whose config-error branches must
|
|
latch discovery_retryable=False (terminal) — and maybe_rediscover must
|
|
INSTALL that terminal config — or a config-invalid IdP (endpoint failing
|
|
SSRF/same-origin) re-probes every cooldown window forever. Drives the
|
|
real discover_oidc: the discovered token_endpoint is on a foreign host,
|
|
so validation rejects it as a config error."""
|
|
from turnstone.core.oidc import maybe_rediscover_oidc
|
|
|
|
state = self._disabled_retryable_state() # issuer=https://idp.example.com
|
|
bad_doc = {
|
|
"authorization_endpoint": "https://idp.example.com/authorize",
|
|
"token_endpoint": "https://attacker.example/token", # foreign host
|
|
"userinfo_endpoint": "https://idp.example.com/userinfo",
|
|
"jwks_uri": "https://idp.example.com/.well-known/jwks.json",
|
|
}
|
|
mock_response = MagicMock()
|
|
mock_response.json.return_value = bad_doc
|
|
mock_response.raise_for_status = MagicMock()
|
|
|
|
probes = {"n": 0}
|
|
|
|
async def _get(url):
|
|
probes["n"] += 1
|
|
return mock_response
|
|
|
|
async def _run():
|
|
client = _mock_async_client(_get)
|
|
with (
|
|
patch("socket.getaddrinfo", return_value=[(2, 1, 6, "", ("93.184.216.34", 0))]),
|
|
patch("httpx.AsyncClient", return_value=client),
|
|
):
|
|
await maybe_rediscover_oidc(state)
|
|
# Same window: cooldown already gates a second probe.
|
|
await maybe_rediscover_oidc(state)
|
|
# Force the cooldown open — but the config is now terminal, so
|
|
# the retryable guard should short-circuit before any probe.
|
|
state.oidc_rediscover_last = None
|
|
await maybe_rediscover_oidc(state)
|
|
|
|
asyncio.run(_run())
|
|
|
|
# Still disabled, but LATCHED terminal (not retryable) — one probe only.
|
|
assert state.oidc_config.enabled is False
|
|
assert state.oidc_config.discovery_retryable is False
|
|
assert probes["n"] == 1
|
|
|
|
def test_rediscover_cooldown_gates_repeat_probes(self):
|
|
from turnstone.core.oidc import maybe_rediscover_oidc
|
|
|
|
state = self._disabled_retryable_state()
|
|
still_down = state.oidc_config # discover keeps returning disabled
|
|
with patch(
|
|
"turnstone.core.oidc.discover_oidc", new=AsyncMock(return_value=still_down)
|
|
) as disc:
|
|
asyncio.run(maybe_rediscover_oidc(state))
|
|
asyncio.run(maybe_rediscover_oidc(state))
|
|
asyncio.run(maybe_rediscover_oidc(state))
|
|
# One IdP probe per cooldown window, however many callers ask.
|
|
disc.assert_awaited_once()
|
|
assert state.oidc_config.enabled is False
|
|
|
|
def test_rediscover_noop_when_not_retryable_or_enabled(self):
|
|
from turnstone.core.oidc import maybe_rediscover_oidc
|
|
|
|
# Operator-disabled (retryable False): never probes.
|
|
state = SimpleNamespace(
|
|
oidc_config=_make_config(
|
|
enabled=False,
|
|
authorization_endpoint="",
|
|
token_endpoint="",
|
|
userinfo_endpoint="",
|
|
jwks_uri="",
|
|
)
|
|
)
|
|
with patch("turnstone.core.oidc.discover_oidc", new=AsyncMock()) as disc:
|
|
asyncio.run(maybe_rediscover_oidc(state))
|
|
disc.assert_not_awaited()
|
|
# Already enabled: never probes.
|
|
state2 = SimpleNamespace(oidc_config=_make_config(enabled=True))
|
|
with patch("turnstone.core.oidc.discover_oidc", new=AsyncMock()) as disc2:
|
|
asyncio.run(maybe_rediscover_oidc(state2))
|
|
disc2.assert_not_awaited()
|
|
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# Lifespan integration: initialize_oidc_state
|
|
# ---------------------------------------------------------------------------
|
|
|
|
|
|
class TestInitializeOIDCState:
|
|
def test_initialize_skips_when_disabled(self):
|
|
"""Disabled config: jwks_data set to None, oidc_config unchanged."""
|
|
cfg = _make_config(enabled=False)
|
|
state = types.SimpleNamespace(oidc_config=cfg, jwks_data="stale")
|
|
|
|
asyncio.run(initialize_oidc_state(state))
|
|
|
|
assert state.oidc_config is cfg
|
|
assert state.jwks_data is None
|
|
|
|
def test_initialize_disables_on_discovery_exception(self):
|
|
"""Unexpected exception in discovery -> enabled flipped to False."""
|
|
cfg = _make_config(
|
|
authorization_endpoint="",
|
|
token_endpoint="",
|
|
jwks_uri="",
|
|
)
|
|
state = types.SimpleNamespace(oidc_config=cfg, jwks_data=None)
|
|
|
|
async def _boom(_cfg):
|
|
raise RuntimeError("boom")
|
|
|
|
with patch("turnstone.core.oidc.discover_oidc", side_effect=_boom):
|
|
asyncio.run(initialize_oidc_state(state))
|
|
|
|
assert state.oidc_config.enabled is False
|
|
assert state.jwks_data is None
|
|
|
|
def test_initialize_disables_on_discovery_returning_disabled(self):
|
|
"""Discovery returns enabled=False (e.g. SSRF reject) -> propagate."""
|
|
cfg = _make_config(
|
|
authorization_endpoint="",
|
|
token_endpoint="",
|
|
jwks_uri="",
|
|
)
|
|
state = types.SimpleNamespace(oidc_config=cfg, jwks_data=None)
|
|
|
|
disabled_cfg = dataclasses.replace(cfg, enabled=False)
|
|
|
|
async def _disabled(_cfg, *, client=None):
|
|
return disabled_cfg
|
|
|
|
with patch("turnstone.core.oidc.discover_oidc", side_effect=_disabled):
|
|
asyncio.run(initialize_oidc_state(state))
|
|
|
|
assert state.oidc_config is disabled_cfg
|
|
assert state.oidc_config.enabled is False
|
|
assert state.jwks_data is None
|
|
|
|
def test_initialize_disables_when_redirect_base_unset(self, caplog):
|
|
"""Discovery succeeds but redirect_base is empty -> disable + log error."""
|
|
cfg = _make_config(redirect_base="")
|
|
state = types.SimpleNamespace(oidc_config=cfg, jwks_data=None)
|
|
|
|
async def _ok(c, *, client=None):
|
|
return c
|
|
|
|
async def _jwks_unexpected(_uri, *, client=None):
|
|
raise AssertionError("fetch_jwks must not be called when redirect_base is empty")
|
|
|
|
with (
|
|
patch("turnstone.core.oidc.discover_oidc", side_effect=_ok),
|
|
patch("turnstone.core.oidc.fetch_jwks", side_effect=_jwks_unexpected),
|
|
caplog.at_level("ERROR", logger="turnstone.core.oidc"),
|
|
):
|
|
asyncio.run(initialize_oidc_state(state))
|
|
|
|
assert state.oidc_config.enabled is False
|
|
assert state.jwks_data is None
|
|
assert any("TURNSTONE_OIDC_REDIRECT_BASE" in record.message for record in caplog.records)
|
|
|
|
def test_initialize_keeps_enabled_but_no_jwks_on_jwks_failure(self):
|
|
"""JWKS fetch failure preserves enabled=True for lazy retry."""
|
|
cfg = _make_config(redirect_base="https://app.example.com")
|
|
state = types.SimpleNamespace(oidc_config=cfg, jwks_data=None)
|
|
|
|
async def _ok(c, *, client=None):
|
|
return c
|
|
|
|
async def _jwks_boom(_uri, *, client=None):
|
|
raise OIDCError("jwks down")
|
|
|
|
with (
|
|
patch("turnstone.core.oidc.discover_oidc", side_effect=_ok),
|
|
patch("turnstone.core.oidc.fetch_jwks", side_effect=_jwks_boom),
|
|
):
|
|
asyncio.run(initialize_oidc_state(state))
|
|
|
|
assert state.oidc_config.enabled is True
|
|
assert state.oidc_config is cfg
|
|
assert state.jwks_data is None
|
|
|
|
def test_initialize_success(self):
|
|
"""Both discovery and JWKS prefetch succeed."""
|
|
cfg = _make_config(redirect_base="https://app.example.com")
|
|
state = types.SimpleNamespace(oidc_config=cfg, jwks_data=None)
|
|
|
|
jwks = {"keys": [{"kid": "k1", "kty": "RSA"}]}
|
|
|
|
async def _ok(c, *, client=None):
|
|
return c
|
|
|
|
async def _jwks(_uri, *, client=None):
|
|
return jwks
|
|
|
|
with (
|
|
patch("turnstone.core.oidc.discover_oidc", side_effect=_ok),
|
|
patch("turnstone.core.oidc.fetch_jwks", side_effect=_jwks),
|
|
):
|
|
asyncio.run(initialize_oidc_state(state))
|
|
|
|
assert state.oidc_config is cfg
|
|
assert state.oidc_config.enabled is True
|
|
assert state.jwks_data == jwks
|
|
assert state.oidc_http_client is not None
|
|
|
|
|
|
async def _async_return(value):
|
|
"""Helper: return a value from an async function."""
|
|
return value
|
|
|
|
|
|
class TestCloseOIDCState:
|
|
def test_close_when_never_initialised(self):
|
|
"""close_oidc_state on bare state must not raise."""
|
|
from turnstone.core.oidc import close_oidc_state
|
|
|
|
state = types.SimpleNamespace()
|
|
asyncio.run(close_oidc_state(state))
|
|
|
|
def test_close_releases_long_lived_client(self):
|
|
"""The long-lived client installed by initialize_oidc_state is aclosed."""
|
|
from turnstone.core.oidc import close_oidc_state
|
|
|
|
client = MagicMock()
|
|
|
|
async def _aclose():
|
|
client.aclose_called = True
|
|
|
|
client.aclose = _aclose
|
|
state = types.SimpleNamespace(oidc_http_client=client)
|
|
asyncio.run(close_oidc_state(state))
|
|
assert getattr(client, "aclose_called", False) is True
|
|
assert state.oidc_http_client is None
|
|
|
|
|
|
class TestLongLivedHTTPClientPassthrough:
|
|
def test_initialize_passes_long_lived_client_to_jwks_only(self):
|
|
"""Discovery uses a transient client; JWKS uses the long-lived one.
|
|
|
|
Splitting the two avoids leaking the long-lived client when a
|
|
disable check (discovery exception, discovery-returned-disabled,
|
|
missing redirect_base) returns before JWKS prefetch — see the
|
|
post-condition contract on initialize_oidc_state.
|
|
"""
|
|
cfg = _make_config(redirect_base="https://app.example.com")
|
|
state = types.SimpleNamespace(oidc_config=cfg, jwks_data=None)
|
|
|
|
seen_clients: list[Any] = []
|
|
|
|
async def _discover(c, *, client=None):
|
|
seen_clients.append(("discover", client))
|
|
return c
|
|
|
|
async def _jwks(_uri, *, client=None):
|
|
seen_clients.append(("jwks", client))
|
|
return {"keys": [{"kid": "k1"}]}
|
|
|
|
with (
|
|
patch("turnstone.core.oidc.discover_oidc", side_effect=_discover),
|
|
patch("turnstone.core.oidc.fetch_jwks", side_effect=_jwks),
|
|
):
|
|
asyncio.run(initialize_oidc_state(state))
|
|
|
|
assert state.oidc_http_client is not None
|
|
kinds = {kind for kind, _ in seen_clients}
|
|
assert {"discover", "jwks"} <= kinds
|
|
|
|
jwks_clients = [c for kind, c in seen_clients if kind == "jwks"]
|
|
assert all(c is state.oidc_http_client for c in jwks_clients)
|
|
# Discovery client is a separate transient (already closed).
|
|
discover_clients = [c for kind, c in seen_clients if kind == "discover"]
|
|
assert all(c is not state.oidc_http_client for c in discover_clients)
|
|
|
|
def test_initialize_does_not_leak_client_on_discovery_failure(self):
|
|
"""Discovery exception path must close the transient and leave http_client None."""
|
|
cfg = _make_config(redirect_base="https://app.example.com")
|
|
state = types.SimpleNamespace(oidc_config=cfg, jwks_data=None)
|
|
|
|
async def _boom(_cfg, *, client=None):
|
|
raise RuntimeError("boom")
|
|
|
|
with patch("turnstone.core.oidc.discover_oidc", side_effect=_boom):
|
|
asyncio.run(initialize_oidc_state(state))
|
|
|
|
assert state.oidc_config.enabled is False
|
|
assert state.oidc_http_client is None
|
|
assert state.jwks_data is None
|
|
|
|
def test_initialize_does_not_leak_client_on_missing_redirect_base(self):
|
|
"""Missing redirect_base path returns before creating the long-lived client."""
|
|
cfg = _make_config(redirect_base="")
|
|
state = types.SimpleNamespace(oidc_config=cfg, jwks_data=None)
|
|
|
|
async def _discover(c, *, client=None):
|
|
return c
|
|
|
|
with patch("turnstone.core.oidc.discover_oidc", side_effect=_discover):
|
|
asyncio.run(initialize_oidc_state(state))
|
|
|
|
assert state.oidc_config.enabled is False
|
|
assert state.oidc_http_client is None
|
|
assert state.jwks_data is None
|