Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
__init__.py147 linesDownload Raw Back to basic_authn
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 
codekingpro/portable-devtools · Team Ai