From 649e44280e0ff6cf41e7b735f87395578a5f3592 Mon Sep 17 00:00:00 2001 From: Manish Kumar Date: Sun, 27 Sep 2026 13:43:22 -0500 Subject: [PATCH] fix: validate Dev Hub JWT claims and forward session token --- api/auth.py | 38 ++++++++--- omnibioai-dev-hub-ui/src/api/client.test.ts | 3 +- omnibioai-dev-hub-ui/src/api/client.ts | 15 ++++- tests/test_auth.py | 75 +++++++++++++++++++-- tests/test_coverage_completion.py | 14 ++-- 5 files changed, 123 insertions(+), 22 deletions(-) diff --git a/api/auth.py b/api/auth.py index 1d6e2be..37278de 100644 --- a/api/auth.py +++ b/api/auth.py @@ -24,6 +24,14 @@ def _jwt_secret() -> str: return os.getenv("JWT_SECRET", "").strip() +def _jwt_audience() -> str: + return os.getenv("JWT_AUDIENCE", "").strip() + + +def _jwt_issuer() -> str: + return os.getenv("JWT_ISSUER", "").strip() + + def validate_auth_config() -> None: """Fail loudly at startup if auth is enabled but misconfigured. @@ -37,11 +45,19 @@ def validate_auth_config() -> None: app's startup path (see api/main.py) -- not from require_auth(), so a misconfigured deployment never serves a single request. """ - if _auth_enabled() and not _jwt_secret(): - raise RuntimeError( - "AUTH_ENABLED is true but JWT_SECRET is not set -- " - "refusing to start with auth silently disabled." - ) + if _auth_enabled(): + missing = [ + name for name, value in ( + ("JWT_SECRET", _jwt_secret()), + ("JWT_AUDIENCE", _jwt_audience()), + ("JWT_ISSUER", _jwt_issuer()), + ) if not value + ] + if missing: + raise RuntimeError( + f"AUTH_ENABLED is true but {', '.join(missing)} is not set -- " + "refusing to start with auth misconfigured." + ) def extract_token(authorization_header: str | None) -> str: @@ -53,9 +69,15 @@ def extract_token(authorization_header: str | None) -> str: return parts[1] -def validate_token(token: str, secret: str) -> dict: +def validate_token(token: str, secret: str, audience: str, issuer: str) -> dict: try: - return _jwt.decode(token, secret, algorithms=["HS256"]) + return _jwt.decode( + token, + secret, + algorithms=["HS256"], + audience=audience, + issuer=issuer, + ) except _jwt.ExpiredSignatureError: raise AuthError("Token has expired", 401) except _jwt.InvalidTokenError as exc: @@ -86,7 +108,7 @@ async def require_auth( try: token = extract_token(authorization) - payload = validate_token(token, internal_secret) + payload = validate_token(token, internal_secret, _jwt_audience(), _jwt_issuer()) for field in ("sub", "service", "username", "email"): value = payload.get(field) if value: diff --git a/omnibioai-dev-hub-ui/src/api/client.test.ts b/omnibioai-dev-hub-ui/src/api/client.test.ts index 810ea93..ebc0979 100644 --- a/omnibioai-dev-hub-ui/src/api/client.test.ts +++ b/omnibioai-dev-hub-ui/src/api/client.test.ts @@ -3,10 +3,12 @@ import { AskError, ASK_ERROR_MESSAGES, describeAskError, getStatus, ragQuery, ra describe("API client", () => { it("sends a JSON query and returns the decoded response", async () => { + document.cookie = "omnibioai_access_token=test-token"; vi.spyOn(globalThis, "fetch").mockResolvedValue(new Response(JSON.stringify({ answer: "ok" }), { status: 200 })); await expect(ragQuery("hello")).resolves.toEqual({ answer: "ok" }); expect(fetch).toHaveBeenCalledWith("/rag/query", expect.objectContaining({ method: "POST", body: JSON.stringify({ query: "hello" }), + headers: { "Content-Type": "application/json", Authorization: "Bearer test-token" }, })); }); @@ -119,4 +121,3 @@ describe("describeAskError", () => { warn.mockRestore(); }); }); - diff --git a/omnibioai-dev-hub-ui/src/api/client.ts b/omnibioai-dev-hub-ui/src/api/client.ts index c803a23..b73471b 100644 --- a/omnibioai-dev-hub-ui/src/api/client.ts +++ b/omnibioai-dev-hub-ui/src/api/client.ts @@ -1,4 +1,15 @@ const API_BASE = ""; +const ACCESS_TOKEN_COOKIE = "omnibioai_access_token"; + +/** Reuse Studio's same-origin session cookie for protected API calls. */ +const authHeaders = (): Record => { + const token = document.cookie + .split(";") + .map((part) => part.trim()) + .find((part) => part.startsWith(`${ACCESS_TOKEN_COOKIE}=`)) + ?.slice(ACCESS_TOKEN_COOKIE.length + 1); + return token ? { Authorization: `Bearer ${decodeURIComponent(token)}` } : {}; +}; // ------------------ ASK OMNIBIOAI ANSWER CONTRACT (ask.v1) ------------------ export interface AskCitation { @@ -80,7 +91,7 @@ export const ragQuery = async (query: string) => { try { res = await fetch(`${API_BASE}/rag/query`, { method: "POST", - headers: { "Content-Type": "application/json" }, + headers: { "Content-Type": "application/json", ...authHeaders() }, body: JSON.stringify({ query }), }); } catch (e) { @@ -102,7 +113,7 @@ export const ragStream = async ( try { const res = await fetch(`${API_BASE}/rag/stream`, { method: "POST", - headers: { "Content-Type": "application/json" }, + headers: { "Content-Type": "application/json", ...authHeaders() }, body: JSON.stringify({ query }), }); diff --git a/tests/test_auth.py b/tests/test_auth.py index c11a0b6..ecf7104 100644 --- a/tests/test_auth.py +++ b/tests/test_auth.py @@ -26,11 +26,25 @@ ) SECRET = "test-secret-value" +AUDIENCE = "omnibioai-platform" +ISSUER = "omnibioai-auth" + + +@pytest.fixture(autouse=True) +def configured_claim_validation(monkeypatch): + monkeypatch.setenv("JWT_AUDIENCE", AUDIENCE) + monkeypatch.setenv("JWT_ISSUER", ISSUER) def _make_token(secret=SECRET, sub="alice", exp_delta=3600, **extra_claims): """Sign an HS256 JWT for tests with the given subject, expiry offset, and extra claims.""" - payload = {"sub": sub, "exp": int(time.time()) + exp_delta, **extra_claims} + payload = { + "sub": sub, + "exp": int(time.time()) + exp_delta, + "aud": AUDIENCE, + "iss": ISSUER, + **extra_claims, + } return jwt.encode(payload, secret, algorithm="HS256") @@ -82,7 +96,7 @@ def test_extract_token_wrong_scheme(): def test_validate_token_valid(): """Return the claims of a correctly signed, unexpired token.""" token = _make_token() - payload = validate_token(token, SECRET) + payload = validate_token(token, SECRET, AUDIENCE, ISSUER) assert payload["sub"] == "alice" @@ -90,20 +104,42 @@ def test_validate_token_expired(): """Reject an expired token with a token-expired AuthError.""" token = _make_token(exp_delta=-3600) # expired an hour ago with pytest.raises(AuthError, match="Token has expired"): - validate_token(token, SECRET) + validate_token(token, SECRET, AUDIENCE, ISSUER) def test_validate_token_malformed_garbage(): """Reject a string that is not a JWT as an invalid token.""" with pytest.raises(AuthError, match="Invalid token"): - validate_token("this-is-not-a-jwt", SECRET) + validate_token("this-is-not-a-jwt", SECRET, AUDIENCE, ISSUER) def test_validate_token_wrong_secret_rejected(): """Reject a token signed with a different secret as invalid.""" token = _make_token(secret=SECRET) with pytest.raises(AuthError, match="Invalid token"): - validate_token(token, "a-different-secret") + validate_token(token, "a-different-secret", AUDIENCE, ISSUER) + + +@pytest.mark.parametrize( + "claims", + [ + {"aud": "wrong-audience"}, + {"aud": None}, + {"iss": "wrong-issuer"}, + {"iss": None}, + ], +) +def test_validate_token_rejects_invalid_or_missing_claims(claims): + extra = {key: value for key, value in claims.items() if value is not None} + if claims.get("aud") is None: + extra["aud"] = None + if claims.get("iss") is None: + extra["iss"] = None + payload = {"sub": "alice", "exp": int(time.time()) + 3600, "aud": AUDIENCE, "iss": ISSUER} + payload.update(extra) + token = jwt.encode(payload, SECRET, algorithm="HS256") + with pytest.raises(AuthError, match="Invalid token"): + validate_token(token, SECRET, AUDIENCE, ISSUER) # ========================================================= @@ -142,6 +178,8 @@ def test_require_auth_enabled_valid_jwt_accepted(monkeypatch): subject.""" monkeypatch.setenv("AUTH_ENABLED", "true") monkeypatch.setenv("JWT_SECRET", SECRET) + monkeypatch.setenv("JWT_AUDIENCE", AUDIENCE) + monkeypatch.setenv("JWT_ISSUER", ISSUER) token = _make_token(sub="alice") resp = client.get("/protected", headers={"Authorization": f"Bearer {token}"}) @@ -154,9 +192,11 @@ def test_require_auth_enabled_jwt_without_identity_claim_falls_back_to_unknown(m """Fall back to the unknown actor when a valid JWT carries no identity claim.""" monkeypatch.setenv("AUTH_ENABLED", "true") monkeypatch.setenv("JWT_SECRET", SECRET) + monkeypatch.setenv("JWT_AUDIENCE", AUDIENCE) + monkeypatch.setenv("JWT_ISSUER", ISSUER) # A structurally valid, correctly-signed token that carries none of the # identity fields require_auth() checks (sub/service/username/email). - payload = {"exp": int(time.time()) + 3600, "role": "irrelevant"} + payload = {"exp": int(time.time()) + 3600, "aud": AUDIENCE, "iss": ISSUER, "role": "irrelevant"} token = jwt.encode(payload, SECRET, algorithm="HS256") resp = client.get("/protected", headers={"Authorization": f"Bearer {token}"}) @@ -169,6 +209,8 @@ def test_require_auth_enabled_expired_jwt_rejected(monkeypatch): """Reject an expired JWT with a 401 that reports the expiry.""" monkeypatch.setenv("AUTH_ENABLED", "true") monkeypatch.setenv("JWT_SECRET", SECRET) + monkeypatch.setenv("JWT_AUDIENCE", AUDIENCE) + monkeypatch.setenv("JWT_ISSUER", ISSUER) token = _make_token(exp_delta=-3600) resp = client.get("/protected", headers={"Authorization": f"Bearer {token}"}) @@ -181,6 +223,8 @@ def test_require_auth_enabled_garbage_token_rejected(monkeypatch): """Reject a malformed Bearer token with a 401 invalid-token error.""" monkeypatch.setenv("AUTH_ENABLED", "true") monkeypatch.setenv("JWT_SECRET", SECRET) + monkeypatch.setenv("JWT_AUDIENCE", AUDIENCE) + monkeypatch.setenv("JWT_ISSUER", ISSUER) resp = client.get("/protected", headers={"Authorization": "Bearer garbage.not.a.jwt"}) @@ -192,6 +236,8 @@ def test_require_auth_enabled_missing_authorization_header_rejected(monkeypatch) """Reject a request with no Authorization header with a 401 when authentication is enabled.""" monkeypatch.setenv("AUTH_ENABLED", "true") monkeypatch.setenv("JWT_SECRET", SECRET) + monkeypatch.setenv("JWT_AUDIENCE", AUDIENCE) + monkeypatch.setenv("JWT_ISSUER", ISSUER) resp = client.get("/protected") @@ -204,6 +250,8 @@ def test_require_auth_internal_header_bypasses_jwt(monkeypatch): JWT.""" monkeypatch.setenv("AUTH_ENABLED", "true") monkeypatch.setenv("JWT_SECRET", SECRET) + monkeypatch.setenv("JWT_AUDIENCE", AUDIENCE) + monkeypatch.setenv("JWT_ISSUER", ISSUER) # No Authorization header at all -- only the matching internal header. resp = client.get("/protected", headers={"X-Devhub-Internal": SECRET}) @@ -216,6 +264,8 @@ def test_require_auth_internal_header_absent_falls_through_to_jwt(monkeypatch): """Fall back to JWT validation when the internal header is absent.""" monkeypatch.setenv("AUTH_ENABLED", "true") monkeypatch.setenv("JWT_SECRET", SECRET) + monkeypatch.setenv("JWT_AUDIENCE", AUDIENCE) + monkeypatch.setenv("JWT_ISSUER", ISSUER) token = _make_token(sub="bob") # No X-Devhub-Internal header sent -- a valid JWT should still work. @@ -229,6 +279,8 @@ def test_require_auth_internal_header_wrong_value_falls_through_and_fails(monkey """Reject a request whose internal header value is wrong and that carries no JWT.""" monkeypatch.setenv("AUTH_ENABLED", "true") monkeypatch.setenv("JWT_SECRET", SECRET) + monkeypatch.setenv("JWT_AUDIENCE", AUDIENCE) + monkeypatch.setenv("JWT_ISSUER", ISSUER) # Wrong internal-header value and no Authorization header -- falls # through to the JWT path, which then has nothing to validate. @@ -278,6 +330,17 @@ def test_validate_auth_config_raises_when_enabled_without_secret(monkeypatch): validate_auth_config() +@pytest.mark.parametrize("missing", ["JWT_AUDIENCE", "JWT_ISSUER"]) +def test_validate_auth_config_requires_claim_configuration(monkeypatch, missing): + monkeypatch.setenv("AUTH_ENABLED", "true") + monkeypatch.setenv("JWT_SECRET", SECRET) + monkeypatch.setenv("JWT_AUDIENCE", AUDIENCE) + monkeypatch.setenv("JWT_ISSUER", ISSUER) + monkeypatch.delenv(missing) + with pytest.raises(RuntimeError, match=missing): + validate_auth_config() + + def test_validate_auth_config_raises_when_enabled_with_blank_secret(monkeypatch): """Treat a whitespace-only JWT secret as unset when authentication is enabled.""" monkeypatch.setenv("AUTH_ENABLED", "true") diff --git a/tests/test_coverage_completion.py b/tests/test_coverage_completion.py index ae139f9..7e2aaee 100644 --- a/tests/test_coverage_completion.py +++ b/tests/test_coverage_completion.py @@ -17,13 +17,15 @@ def test_auth_token_helpers_cover_success_and_failures(monkeypatch): with pytest.raises(auth.AuthError, match="Bearer"): auth.extract_token("Token abc") - token = jwt.encode({"sub": "user-1"}, "secret", algorithm="HS256") - assert auth.validate_token(token, "secret")["sub"] == "user-1" + token = jwt.encode({"sub": "user-1", "aud": "audience", "iss": "issuer", "exp": 4102444800}, "secret", algorithm="HS256") + assert auth.validate_token(token, "secret", "audience", "issuer")["sub"] == "user-1" with pytest.raises(auth.AuthError, match="Invalid token"): - auth.validate_token("bad", "secret") + auth.validate_token("bad", "secret", "audience", "issuer") monkeypatch.setenv("AUTH_ENABLED", "true") monkeypatch.setenv("JWT_SECRET", "secret") + monkeypatch.setenv("JWT_AUDIENCE", "audience") + monkeypatch.setenv("JWT_ISSUER", "issuer") assert auth._auth_enabled() is True assert auth._jwt_secret() == "secret" @@ -35,12 +37,14 @@ async def test_require_auth_all_modes_and_identity_fields(monkeypatch): monkeypatch.setenv("AUTH_ENABLED", "true") monkeypatch.setenv("JWT_SECRET", "secret") + monkeypatch.setenv("JWT_AUDIENCE", "audience") + monkeypatch.setenv("JWT_ISSUER", "issuer") assert await auth.require_auth(SimpleNamespace(), x_devhub_internal="secret") == "devhub-ui" - token = jwt.encode({"email": "user@example.com"}, "secret", algorithm="HS256") + token = jwt.encode({"email": "user@example.com", "aud": "audience", "iss": "issuer", "exp": 4102444800}, "secret", algorithm="HS256") assert await auth.require_auth(SimpleNamespace(), authorization=f"Bearer {token}") == "user@example.com" - unknown = jwt.encode({}, "secret", algorithm="HS256") + unknown = jwt.encode({"aud": "audience", "iss": "issuer", "exp": 4102444800}, "secret", algorithm="HS256") assert await auth.require_auth(SimpleNamespace(), authorization=f"Bearer {unknown}") == "unknown" with pytest.raises(HTTPException, match="Authorization header"):