codekingpro/portable-devtools
115k
1from typing import Any, Dict, Mapping, Optional, TypeVar
2from urllib.parse import quote, urlparse, urlunparse
3import logging
4import orjson as json
5import httpx
6
7import chromadb.errors as errors
8from chromadb.config import Component, Settings, System
9
10logger = logging.getLogger(__name__)
11
12
13# inherits from Component so that it can create an init function to use system
14# this way it can build limits from the settings in System
15class BaseHTTPClient(Component):
16 _settings: Settings
17 pre_flight_checks: Any = None
18 DEFAULT_KEEPALIVE_SECS: float = 40.0
19
20 def __init__(self, system: System):
21 super().__init__(system)
22 self._settings = system.settings
23 keepalive_setting = self._settings.chroma_http_keepalive_secs
24 self.keepalive_secs: Optional[float] = (
25 keepalive_setting
26 if keepalive_setting is not None
27 else BaseHTTPClient.DEFAULT_KEEPALIVE_SECS
28 )
29 self._http_limits = self._build_limits()
30
31 def _build_limits(self) -> httpx.Limits:
32 limit_kwargs: Dict[str, Any] = {}
33 if self.keepalive_secs is not None:
34 limit_kwargs["keepalive_expiry"] = self.keepalive_secs
35
36 max_connections = self._settings.chroma_http_max_connections
37 if max_connections is not None:
38 limit_kwargs["max_connections"] = max_connections
39
40 max_keepalive_connections = self._settings.chroma_http_max_keepalive_connections
41 if max_keepalive_connections is not None:
42 limit_kwargs["max_keepalive_connections"] = max_keepalive_connections
43
44 return httpx.Limits(**limit_kwargs)
45
46 @property
47 def http_limits(self) -> httpx.Limits:
48 return self._http_limits
49
50 @staticmethod
51 def _validate_host(host: str) -> None:
52 parsed = urlparse(host)
53 if "/" in host and parsed.scheme not in {"http", "https"}:
54 raise ValueError(
55 "Invalid URL. " f"Unrecognized protocol - {parsed.scheme}."
56 )
57 if "/" in host and (not host.startswith("http")):
58 raise ValueError(
59 "Invalid URL. "
60 "Seems that you are trying to pass URL as a host but without \
61 specifying the protocol. "
62 "Please add http:// or https:// to the host."
63 )
64
65 @staticmethod
66 def resolve_url(
67 chroma_server_host: str,
68 chroma_server_ssl_enabled: Optional[bool] = False,
69 default_api_path: Optional[str] = "",
70 chroma_server_http_port: Optional[int] = 8000,
71 ) -> str:
72 _skip_port = False
73 _chroma_server_host = chroma_server_host
74 BaseHTTPClient._validate_host(_chroma_server_host)
75 if _chroma_server_host.startswith("http"):
76 logger.debug("Skipping port as the user is passing a full URL")
77 _skip_port = True
78 parsed = urlparse(_chroma_server_host)
79
80 scheme = "https" if chroma_server_ssl_enabled else parsed.scheme or "http"
81 net_loc = parsed.netloc or parsed.hostname or chroma_server_host
82 port = (
83 ":" + str(parsed.port or chroma_server_http_port) if not _skip_port else ""
84 )
85 path = parsed.path or default_api_path
86
87 if not path or path == net_loc:
88 path = default_api_path if default_api_path else ""
89 if not path.endswith(default_api_path or ""):
90 path = path + default_api_path if default_api_path else ""
91 full_url = urlunparse(
92 (scheme, f"{net_loc}{port}", quote(path.replace("//", "/")), "", "", "")
93 )
94
95 return full_url
96
97 # requests removes None values from the built query string, but httpx includes it as an empty value
98 T = TypeVar("T", bound=Dict[Any, Any])
99
100 @staticmethod
101 def _clean_params(params: T) -> T:
102 """Remove None values from provided dict."""
103 return {k: v for k, v in params.items() if v is not None} # type: ignore
104
105 @staticmethod
106 def _raise_chroma_error(resp: httpx.Response) -> None:
107 """Raises an error if the response is not ok, using a ChromaError if possible."""
108 try:
109 resp.raise_for_status()
110 return
111 except httpx.HTTPStatusError:
112 pass
113
114 chroma_error = None
115 try:
116 body = json.loads(resp.text)
117 if "error" in body:
118 if body["error"] in errors.error_types:
119 chroma_error = errors.error_types[body["error"]](body["message"])
120
121 trace_id = resp.headers.get("chroma-trace-id")
122 if trace_id:
123 chroma_error.trace_id = trace_id
124
125 except BaseException:
126 pass
127
128 if chroma_error:
129 raise chroma_error
130
131 try:
132 resp.raise_for_status()
133 except httpx.HTTPStatusError:
134 trace_id = resp.headers.get("chroma-trace-id")
135 if trace_id:
136 raise Exception(f"{resp.text} (trace ID: {trace_id})")
137 raise (Exception(resp.text))
138
139 def get_request_headers(self) -> Mapping[str, str]:
140 """Return headers used for HTTP requests."""
141 return {}
142
143 def get_api_url(self) -> str:
144 """Return the API URL for this client."""
145 return ""
146 