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
11 changes: 8 additions & 3 deletions app/api/routes_apikeys.py
Original file line number Diff line number Diff line change
Expand Up @@ -29,6 +29,7 @@ def _key_out(key: ApiKey) -> ApiKeyOut:
created_at=key.created_at,
expires_at=key.expires_at,
last_used_at=key.last_used_at,
test=apikey_service.is_test_key(key),
)


Expand All @@ -43,7 +44,7 @@ def create_api_key(
try:
api_key, full_key = apikey_service.create_api_key(
db, org_id, membership.user_id, body.name, body.scopes, caller_permissions,
expires_at=body.expires_at,
expires_at=body.expires_at, test=body.test,
)
except ValueError as e:
raise HTTPException(400, str(e))
Expand All @@ -53,6 +54,7 @@ def create_api_key(
key_prefix=api_key.key_prefix,
scopes=api_key.scopes or [],
expires_at=api_key.expires_at,
test=body.test,
key=full_key,
)

Expand Down Expand Up @@ -114,7 +116,8 @@ def exchange_api_key(
db: Session = Depends(get_db),
x_api_key_exchange_secret: str = Header(default=""),
):
"""Gateway-only: trade an omni_sk_ key for a short-lived access token.
"""Gateway-only: trade an omni_sk_live_/omni_sk_test_ (or a pre-M13
bare omni_sk_) key for a short-lived access token.

Callable only with the shared API_KEY_EXCHANGE_SECRET, so the minted
token can't be obtained by a key holder directly and used to reach
Expand Down Expand Up @@ -210,6 +213,7 @@ def _me_key_out(key: ApiKey) -> ApiKeyOut:
created_at=key.created_at,
expires_at=key.expires_at,
last_used_at=key.last_used_at,
test=apikey_service.is_test_key(key),
)


Expand Down Expand Up @@ -239,7 +243,7 @@ def create_my_api_key(
api_key, full_key = apikey_service.create_api_key(
db, membership.organization_id, membership.user_id, body.name, internal_scopes,
org_service.permissions_for_membership(membership),
expires_at=body.expires_at,
expires_at=body.expires_at, test=body.test,
)
except ValueError as e:
raise HTTPException(400, str(e))
Expand All @@ -249,6 +253,7 @@ def create_my_api_key(
key_prefix=api_key.key_prefix,
scopes=_to_public_scopes(api_key.scopes or []),
expires_at=api_key.expires_at,
test=body.test,
key=full_key,
)

Expand Down
4 changes: 4 additions & 0 deletions app/schemas/apikeys.py
Original file line number Diff line number Diff line change
Expand Up @@ -7,6 +7,7 @@ class ApiKeyCreate(BaseModel):
name: str
scopes: list[str] = []
expires_at: datetime | None = None # M9: optional, self-service-settable; None = no expiry
test: bool = False # M13: omni_sk_test_ instead of omni_sk_live_; see apikey_service.is_test_key


class ApiKeyRename(BaseModel):
Expand All @@ -19,6 +20,7 @@ class ApiKeyCreated(BaseModel):
key_prefix: str
scopes: list[str]
expires_at: datetime | None = None
test: bool = False
key: str # full plaintext key -- returned exactly once, at creation


Expand All @@ -31,6 +33,7 @@ class ApiKeyOut(BaseModel):
created_at: datetime | None
expires_at: datetime | None
last_used_at: datetime | None
test: bool = False


class ApiKeyExchangeIn(BaseModel):
Expand All @@ -44,3 +47,4 @@ class ApiKeyExchangeOut(BaseModel):
organization_id: int
user_id: int
permissions: list[str]
test_mode: bool = False
35 changes: 31 additions & 4 deletions app/services/apikey_service.py
Original file line number Diff line number Diff line change
Expand Up @@ -13,11 +13,32 @@

_ALPHABET = "abcdefghijklmnopqrstuvwxyzABCDEFGHIJKLMNOPQRSTUVWXYZ0123456789"
_SECRET_LENGTH = 40
# M13 (design audit gap #9): omni_sk_live_/omni_sk_test_ replace the single
# omni_sk_ prefix every key issued before this milestone still uses. Both
# new prefixes start with the old bare one, so is_test_key()/exchange_api_key
# below stay the only two places that need to know the difference --
# every pre-M13 key is unambiguously "live" (it predates test mode
# existing at all), and nothing else anywhere needs to distinguish the
# three shapes. _PREFIX stays the "is this an API key at all" check (also
# used, unchanged, by omnibioai-api-gateway's own is_api_key()).
_PREFIX = "omni_sk_"
_LIVE_PREFIX = "omni_sk_live_"
_TEST_PREFIX = "omni_sk_test_"


def _generate_key() -> str:
return _PREFIX + "".join(secrets.choice(_ALPHABET) for _ in range(_SECRET_LENGTH))
def _generate_key(test: bool = False) -> str:
prefix = _TEST_PREFIX if test else _LIVE_PREFIX
return prefix + "".join(secrets.choice(_ALPHABET) for _ in range(_SECRET_LENGTH))


def is_test_key(api_key: ApiKey) -> bool:
"""Whether `api_key` is a test-mode key (omni_sk_test_) -- derived from
the already-stored, non-secret key_prefix rather than a separate DB
column, since that prefix already encodes the answer. A pre-M13 key
(bare omni_sk_ prefix) is always False: it predates test mode, so it
is unambiguously live, the same posture a pre-M9 key's scopes get
read with (see routes_apikeys.py's _to_public_scopes fallback)."""
return bool(api_key.key_prefix) and api_key.key_prefix.startswith(_TEST_PREFIX)


def _hash_key(full_key: str) -> str:
Expand Down Expand Up @@ -45,6 +66,7 @@ def create_api_key(
scopes: list[str],
caller_permissions: set[str],
expires_at: datetime | None = None,
test: bool = False,
) -> tuple[ApiKey, str]:
"""Returns (ApiKey row, full plaintext key). The plaintext is never
persisted -- only its sha256 hash is stored -- so this is the only
Expand All @@ -63,12 +85,13 @@ def create_api_key(
if expires_at is not None and expires_at <= datetime.utcnow():
raise ValueError("expires_at must be in the future")

full_key = _generate_key()
full_key = _generate_key(test=test)
prefix_len = len(_TEST_PREFIX if test else _LIVE_PREFIX)
api_key = ApiKey(
organization_id=organization_id,
created_by_user_id=creator_user_id,
name=name,
key_prefix=full_key[: len(_PREFIX) + 4],
key_prefix=full_key[: prefix_len + 4],
key_hash=_hash_key(full_key),
scopes=scopes,
status="active",
Expand Down Expand Up @@ -242,4 +265,8 @@ def exchange_api_key(db: Session, full_key: str) -> dict | None:
"organization_id": api_key.organization_id,
"user_id": user.id,
"permissions": permissions,
# M13: the gateway reads this to serve a canned, unbilled response
# instead of forwarding to RAG/consuming real quota -- see
# omnibioai-api-gateway's app/routes/v1.py.
"test_mode": is_test_key(api_key),
}
65 changes: 65 additions & 0 deletions tests/test_apikey_exchange.py
Original file line number Diff line number Diff line change
Expand Up @@ -165,6 +165,71 @@ def test_exchange_rejects_key_of_suspended_issuer(client, org_key, exchange_secr
assert _exchange(client, org_key["key"]).status_code == 401


def test_exchange_reports_test_mode_false_for_a_live_key(client, org_key, exchange_secret):
"""org_key's fixture creates a key without test=True -- the default,
omni_sk_live_, must report test_mode False."""
data = _exchange(client, org_key["key"]).json()
assert data["test_mode"] is False


def test_exchange_reports_test_mode_true_for_a_test_key(client, exchange_secret):
owner = _register_and_login(client)
headers = _auth_header(owner["access_token"])
org = client.post(
"/orgs", json={"name": "Test Mode Org", "slug": f"test-mode-{uuid.uuid4().hex[:8]}"}, headers=headers,
).json()
created = client.post(
f"/orgs/{org['id']}/api-keys",
json={"name": "sandbox", "scopes": [], "test": True},
headers=headers,
).json()
assert created["key"].startswith("omni_sk_test_")
assert created["test"] is True

data = _exchange(client, created["key"]).json()
assert data["test_mode"] is True


def test_exchange_reports_test_mode_false_for_a_pre_m13_legacy_key(client, exchange_secret):
"""A key issued before M13 has the bare omni_sk_ prefix (not
omni_sk_live_/omni_sk_test_) -- it predates test mode entirely, so it
must still exchange successfully and report test_mode False, not
error out on an unrecognized prefix shape."""
import hashlib
from datetime import datetime

owner = _register_and_login(client)
headers = _auth_header(owner["access_token"])
org = client.post(
"/orgs", json={"name": "Legacy Org", "slug": f"legacy-{uuid.uuid4().hex[:8]}"}, headers=headers,
).json()

legacy_key = "omni_sk_" + "z" * 40
db = _DirectSession()
try:
membership = (
db.query(OrganizationMembership)
.filter(OrganizationMembership.organization_id == org["id"])
.one()
)
db.add(ApiKey(
organization_id=org["id"],
created_by_user_id=membership.user_id,
name="legacy",
key_prefix=legacy_key[:12],
key_hash=hashlib.sha256(legacy_key.encode()).hexdigest(),
scopes=[],
status="active",
created_at=datetime.utcnow(),
))
db.commit()
finally:
db.close()

data = _exchange(client, legacy_key).json()
assert data["test_mode"] is False


def test_exchange_throttles_last_used_writes(client, org_key, exchange_secret):
def last_used():
db = _DirectSession()
Expand Down
24 changes: 24 additions & 0 deletions tests/test_apikeys.py
Original file line number Diff line number Diff line change
Expand Up @@ -142,6 +142,30 @@ def test_create_api_key_rejects_past_expires_at(client, org):
assert resp.status_code == 400


def test_create_api_key_defaults_to_live_mode(client, org):
resp = client.post(
f"/orgs/{org['id']}/api-keys", json={"name": "Default mode", "scopes": []}, headers=org["owner_headers"],
)
assert resp.status_code == 201
data = resp.json()
assert data["key"].startswith("omni_sk_live_")
assert data["test"] is False


def test_create_api_key_with_test_true_issues_a_test_key(client, org):
resp = client.post(
f"/orgs/{org['id']}/api-keys", json={"name": "Sandbox", "scopes": [], "test": True},
headers=org["owner_headers"],
)
assert resp.status_code == 201
data = resp.json()
assert data["key"].startswith("omni_sk_test_")
assert data["test"] is True

listed = client.get(f"/orgs/{org['id']}/api-keys", headers=org["owner_headers"]).json()
assert next(k for k in listed if k["id"] == data["id"])["test"] is True


def test_rename_api_key(client, org):
"""PATCH renames a key and the new name is reflected in a subsequent listing."""
create = client.post(
Expand Down
18 changes: 18 additions & 0 deletions tests/test_me_api_keys.py
Original file line number Diff line number Diff line change
Expand Up @@ -100,6 +100,24 @@ def test_create_rejects_past_expires_at(client, scientist):
assert resp.status_code == 400


def test_create_defaults_to_live_mode(client, scientist):
resp = client.post("/me/api-keys", json={"name": "notebook"}, headers=scientist["headers"])
created = resp.json()
assert created["key"].startswith("omni_sk_live_")
assert created["test"] is False


def test_create_with_test_true_issues_a_test_key(client, scientist):
resp = client.post("/me/api-keys", json={"name": "sandbox", "test": True}, headers=scientist["headers"])
assert resp.status_code == 201
created = resp.json()
assert created["key"].startswith("omni_sk_test_")
assert created["test"] is True

listed = client.get("/me/api-keys", headers=scientist["headers"]).json()
assert next(k for k in listed if k["id"] == created["id"])["test"] is True


def test_rename_own_key(client, scientist):
created = client.post("/me/api-keys", json={"name": "old"}, headers=scientist["headers"]).json()

Expand Down
Loading