codekingpro/portable-devtools
114k
1import base64
2import random
3import re
4import time
5import traceback
6
7import bcrypt
8import logging
9
10from overrides import override
11from pydantic import SecretStr
12
13from chromadb.auth import (
14 UserIdentity,
15 ServerAuthenticationProvider,
16 ClientAuthProvider,
17 ClientAuthHeaders,
18 AuthError,
19)
20from chromadb.config import System
21from chromadb.errors import ChromaAuthError
22from chromadb.telemetry.opentelemetry import (
23 OpenTelemetryGranularity,
24 trace_method,
25)
26
27
28from typing import Dict
29
30
31logger = logging.getLogger(__name__)
32
33__all__ = ["BasicAuthenticationServerProvider", "BasicAuthClientProvider"]
34
35AUTHORIZATION_HEADER = "Authorization"
36
37
38class BasicAuthClientProvider(ClientAuthProvider):
39 """
40 Client auth provider for basic auth. The credentials are passed as a
41 base64-encoded string in the Authorization header prepended with "Basic ".
42 """
43
44 def __init__(self, system: System) -> None:
45 super().__init__(system)
46 self._settings = system.settings
47 system.settings.require("chroma_client_auth_credentials")
48 self._creds = SecretStr(str(system.settings.chroma_client_auth_credentials))
49
50 @override
51 def authenticate(self) -> ClientAuthHeaders:
52 encoded = base64.b64encode(
53 f"{self._creds.get_secret_value()}".encode("utf-8")
54 ).decode("utf-8")
55 return {
56 AUTHORIZATION_HEADER: SecretStr(f"Basic {encoded}"),
57 }
58
59
60class BasicAuthenticationServerProvider(ServerAuthenticationProvider):
61 """
62 Server auth provider for basic auth. The credentials are read from
63 `chroma_server_authn_credentials_file` and each line must be in the format
64 <username>:<bcrypt passwd>.
65
66 Expects tokens to be passed as a base64-encoded string in the Authorization
67 header prepended with "Basic".
68 """
69
70 def __init__(self, system: System) -> None:
71 super().__init__(system)
72 self._settings = system.settings
73
74 self._creds: Dict[str, SecretStr] = {}
75 creds = self.read_creds_or_creds_file()
76
77 for line in creds:
78 if not line.strip():
79 continue
80 _raw_creds = [v for v in line.strip().split(":")]
81 if (
82 _raw_creds
83 and _raw_creds[0]
84 and len(_raw_creds) != 2
85 or not all(_raw_creds)
86 ):
87 raise ValueError(
88 f"Invalid htpasswd credentials found: {_raw_creds}. "
89 "Lines must be exactly <username>:<bcrypt passwd>."
90 )
91 username = _raw_creds[0]
92 password = _raw_creds[1]
93 if username in self._creds:
94 raise ValueError(
95 "Duplicate username found in "
96 "[chroma_server_authn_credentials]. "
97 "Usernames must be unique."
98 )
99 self._creds[username] = SecretStr(password)
100
101 @trace_method(
102 "BasicAuthenticationServerProvider.authenticate", OpenTelemetryGranularity.ALL
103 )
104 @override
105 def authenticate_or_raise(self, headers: Dict[str, str]) -> UserIdentity:
106 try:
107 if AUTHORIZATION_HEADER.lower() not in headers.keys():
108 raise AuthError(AUTHORIZATION_HEADER + " header not found")
109 _auth_header = headers[AUTHORIZATION_HEADER.lower()]
110 _auth_header = re.sub(r"^Basic ", "", _auth_header)
111 _auth_header = _auth_header.strip()
112
113 base64_decoded = base64.b64decode(_auth_header).decode("utf-8")
114 if ":" not in base64_decoded:
115 raise AuthError("Invalid Authorization header format")
116 username, password = base64_decoded.split(":", 1)
117 username = str(username) # convert to string to prevent header injection
118 password = str(password) # convert to string to prevent header injection
119 if username not in self._creds:
120 raise AuthError("Invalid username or password")
121
122 _pwd_check = bcrypt.checkpw(
123 password.encode("utf-8"),
124 self._creds[username].get_secret_value().encode("utf-8"),
125 )
126 if not _pwd_check:
127 raise AuthError("Invalid username or password")
128 return UserIdentity(user_id=username)
129 except AuthError as e:
130 logger.error(
131 f"BasicAuthenticationServerProvider.authenticate failed: {repr(e)}"
132 )
133 except Exception as e:
134 tb = traceback.extract_tb(e.__traceback__)
135 # Get the last call stack
136 last_call_stack = tb[-1]
137 line_number = last_call_stack.lineno
138 filename = last_call_stack.filename
139 logger.error(
140 "BasicAuthenticationServerProvider.authenticate failed: "
141 f"Failed to authenticate {type(e).__name__} at {filename}:{line_number}"
142 )
143 time.sleep(
144 random.uniform(0.001, 0.005)
145 ) # add some jitter to avoid timing attacks
146 raise ChromaAuthError()
147 