Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
38 changes: 30 additions & 8 deletions api/auth.py
Original file line number Diff line number Diff line change
Expand Up @@ -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.

Expand All @@ -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:
Expand All @@ -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:
Expand Down Expand Up @@ -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:
Expand Down
3 changes: 2 additions & 1 deletion omnibioai-dev-hub-ui/src/api/client.test.ts
Original file line number Diff line number Diff line change
Expand Up @@ -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" },
}));
});

Expand Down Expand Up @@ -119,4 +121,3 @@ describe("describeAskError", () => {
warn.mockRestore();
});
});

15 changes: 13 additions & 2 deletions omnibioai-dev-hub-ui/src/api/client.ts
Original file line number Diff line number Diff line change
@@ -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<string, string> => {
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 {
Expand Down Expand Up @@ -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) {
Expand All @@ -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 }),
});

Expand Down
75 changes: 69 additions & 6 deletions tests/test_auth.py
Original file line number Diff line number Diff line change
Expand Up @@ -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")


Expand Down Expand Up @@ -82,28 +96,50 @@ 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"


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)


# =========================================================
Expand Down Expand Up @@ -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}"})
Expand All @@ -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}"})
Expand All @@ -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}"})
Expand All @@ -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"})

Expand All @@ -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")

Expand All @@ -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})
Expand All @@ -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.
Expand All @@ -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.
Expand Down Expand Up @@ -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")
Expand Down
14 changes: 9 additions & 5 deletions tests/test_coverage_completion.py
Original file line number Diff line number Diff line change
Expand Up @@ -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"

Expand All @@ -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"):
Expand Down
Loading