Coverage for src/lexigram/auth/authn/relay.py: 96%

54 statements  

« prev     ^ index     » next       coverage.py v7.15.4, created at 2026-08-26 00:58 +0800

1"""Relay gateway API-key verifier (new-api TokenAuth parity). 

2 

3Application-level adapter binding lexigram-auth's ``APIKeyManager`` 

4behind the framework's ``RelayAuthVerifierProtocol``, so the relay 

5gateway can enforce API-key auth on its inbound routes without ever 

6importing lexigram-auth. Credential parsing follows new-api's 

7``TokenAuth`` middleware: Claude routes read ``x-api-key``, Gemini 

8routes read ``?key=`` then ``x-goog-api-key``, everything else falls 

9back to the ``Authorization`` bearer header; keys are trimmed of the 

10``sk-`` family prefix before validation. 

11""" 

12 

13from __future__ import annotations 

14 

15from collections.abc import Awaitable, Callable 

16 

17from lexigram.auth.authn.apikeys import APIKeyManager 

18from lexigram.contracts.ai.relay import ( 

19 RelayAuthError, 

20 RelayAuthIdentity, 

21) 

22from lexigram.contracts.core.result import Err, Ok, Result 

23from lexigram.logging import get_logger 

24 

25logger = get_logger(__name__) 

26 

27__all__ = ["RelayApiKeyVerifier"] 

28 

29 

30class RelayApiKeyVerifier: 

31 """Verify inbound relay keys through lexigram-auth's API key manager. 

32 

33 The verifier parses credentials from the request in new-api's 

34 precedence order (path-scoped ``x-api-key`` / ``?key=`` / 

35 ``x-goog-api-key`` first, ``Authorization`` bearer last), normalizes 

36 the key, and delegates lookup to an injected 

37 :class:`~lexigram.auth.authn.apikeys.APIKeyManager`. It never logs 

38 key material and never raises for a bad credential. 

39 

40 Args: 

41 manager: The API key manager performing validation and hashing. 

42 user_status_checker: Optional async predicate over a key owner's 

43 user id; when provided the owner must pass it or the request 

44 is rejected as ``AUTH_USER_DISABLED``. When ``None`` only 

45 key validity governs. 

46 """ 

47 

48 def __init__( 

49 self, 

50 manager: APIKeyManager, 

51 *, 

52 user_status_checker: Callable[[str], Awaitable[bool]] | None = None, 

53 ) -> None: 

54 """Initialise the verifier with the key manager and checks.""" 

55 self._manager = manager 

56 self._user_status_checker = user_status_checker 

57 

58 async def authenticate( 

59 self, request: object 

60 ) -> Result[RelayAuthIdentity, RelayAuthError]: 

61 """Verify the caller of an inbound relay request. 

62 

63 Args: 

64 request: The inbound request; duck-typed for ``headers``, 

65 ``query_params``, ``path``, and ``client``. 

66 

67 Returns: 

68 ``Ok(identity)`` with the key's user, token, and prefix on 

69 success; ``Err`` with ``AUTH_TOKEN_INVALID`` for missing or 

70 unknown keys and ``AUTH_USER_DISABLED`` for disabled owners. 

71 """ 

72 headers = _as_dict(getattr(request, "headers", None)) 

73 query = _as_dict(getattr(request, "query_params", None)) 

74 path: str = getattr(request, "path", "") or "" 

75 

76 raw_key = self._extract_key(headers, query, path) 

77 if not raw_key: 

78 return Err(RelayAuthError("AUTH_TOKEN_INVALID", "missing API key")) 

79 

80 key = _normalize(raw_key) 

81 if not key: 

82 return Err(RelayAuthError("AUTH_TOKEN_INVALID", "missing API key")) 

83 

84 client = getattr(request, "client", None) 

85 ip_address = getattr(client, "host", None) if client else None 

86 

87 api_key = await self._manager.validate_key(key, ip_address=ip_address) 

88 if api_key is None: 

89 logger.warning("relay_auth_invalid_key", key_prefix=key[:8]) 

90 return Err(RelayAuthError("AUTH_TOKEN_INVALID", "invalid API key")) 

91 

92 if ( 

93 self._user_status_checker is not None 

94 and not await self._user_status_checker(api_key.user_id) 

95 ): 

96 return Err(RelayAuthError("AUTH_USER_DISABLED", "user is disabled")) 

97 

98 return Ok( 

99 RelayAuthIdentity( 

100 user_id=api_key.user_id, 

101 token_id=api_key.key_id, 

102 key_prefix=api_key.prefix, 

103 ) 

104 ) 

105 

106 def _extract_key( 

107 self, headers: dict[str, str], query: dict[str, str], path: str 

108 ) -> str: 

109 """Return the raw key per new-api's precedence order.""" 

110 if "/v1/messages" in path: 

111 return headers.get("x-api-key", "") or "" 

112 

113 if path.startswith(("/v1beta/models", "/v1beta/openai/models", "/v1/models/")): 

114 key = query.get("key", "") 

115 if key: 

116 return key 

117 return headers.get("x-goog-api-key", "") or "" 

118 

119 authorization = headers.get("authorization", "") 

120 if authorization.startswith(("Bearer ", "bearer ")): 

121 return authorization[7:].strip() 

122 return "" 

123 

124 

125def _as_dict(value: object) -> dict[str, str]: 

126 """Coerce a header/query mapping into a plain case-folded dict. 

127 

128 Starlette headers and query params are case-insensitive mappings; 

129 extracting via a casefolded copy mirrors that behavior for the 

130 duck-typed request doubles used in tests. 

131 """ 

132 pairs: dict[str, str] = {} 

133 for key, item in getattr(value, "items", list)() if value else []: 

134 if isinstance(key, str) and isinstance(item, str): 

135 pairs[key.casefold()] = item 

136 return pairs 

137 

138 

139def _normalize(raw_key: str) -> str: 

140 """Trim the ``sk-`` family prefix, keeping the first dash segment. 

141 

142 Harmless for lexigram's ``sk_live_...`` keys (no leading dash) and 

143 makes OpenAI-style ``sk-abc-123`` keys validate against their first 

144 segment, mirroring new-api's ``TokenAuth``. 

145 """ 

146 key = raw_key 

147 if key.startswith("sk-"): 

148 key = key[3:] 

149 return key.split("-", 1)[0].strip()