codekingpro/portable-devtools
114k
1from typing import (
2 Any,
3 Callable,
4 cast,
5 Dict,
6 Sequence,
7 Optional,
8 Type,
9 TypeVar,
10 Tuple,
11)
12import fastapi
13import orjson
14from anyio import (
15 to_thread,
16 CapacityLimiter,
17)
18from fastapi import FastAPI as _FastAPI, Response, Request
19from fastapi.openapi.utils import get_openapi
20from fastapi.middleware.cors import CORSMiddleware
21from fastapi.responses import ORJSONResponse
22from fastapi.routing import APIRoute
23from fastapi import HTTPException, status
24from functools import wraps
25
26from chromadb.api.collection_configuration import (
27 load_update_collection_configuration_from_json,
28 load_create_collection_configuration_from_json,
29 create_collection_configuration_from_legacy_collection_metadata,
30 CreateCollectionConfiguration,
31)
32from pydantic import BaseModel
33from chromadb import __version__ as chromadb_version
34from chromadb.api.types import (
35 DeleteResult,
36 Embedding,
37 GetResult,
38 QueryResult,
39 Embeddings,
40 convert_list_embeddings_to_np,
41)
42from chromadb.auth import UserIdentity
43from chromadb.auth import (
44 AuthzAction,
45 AuthzResource,
46 ServerAuthenticationProvider,
47 ServerAuthorizationProvider,
48)
49from chromadb.config import DEFAULT_DATABASE, DEFAULT_TENANT, Settings, System
50from chromadb.api import ServerAPI
51from chromadb.errors import (
52 ChromaError,
53 InvalidDimensionException,
54 InvalidHTTPVersion,
55 RateLimitError,
56 QuotaError,
57)
58from chromadb.quota import QuotaEnforcer
59from chromadb.rate_limit import AsyncRateLimitEnforcer
60from chromadb.server import Server
61from chromadb.server.fastapi.types import (
62 AddEmbedding,
63 CreateDatabase,
64 CreateTenant,
65 DeleteEmbedding,
66 GetEmbedding,
67 QueryEmbedding,
68 CreateCollection,
69 UpdateCollection,
70 UpdateEmbedding,
71)
72from starlette.datastructures import Headers
73import logging
74
75from chromadb.telemetry.product.events import ServerStartEvent
76from chromadb.utils.fastapi import fastapi_json_response, string_to_uuid as _uuid
77from opentelemetry import trace
78
79from chromadb.telemetry.opentelemetry.fastapi import instrument_fastapi
80from chromadb.types import Database, Tenant
81from chromadb.telemetry.product import ServerContext, ProductTelemetryClient
82from chromadb.telemetry.opentelemetry import (
83 OpenTelemetryClient,
84 OpenTelemetryGranularity,
85 add_attributes_to_current_span,
86 trace_method,
87)
88from chromadb.types import Collection as CollectionModel
89
90logger = logging.getLogger(__name__)
91
92
93def rate_limit(func):
94 @wraps(func)
95 async def wrapper(*args: Any, **kwargs: Any) -> Any:
96 self = args[0]
97 return await self._async_rate_limit_enforcer.rate_limit(func)(*args, **kwargs)
98
99 return wrapper
100
101
102def use_route_names_as_operation_ids(app: _FastAPI) -> None:
103 """
104 Simplify operation IDs so that generated API clients have simpler function
105 names.
106 Should be called only after all routes have been added.
107 """
108 for route in app.routes:
109 if isinstance(route, APIRoute):
110 route.operation_id = route.name + ("-v2" if "v2" in route.path else "-v1")
111
112
113async def add_trace_id_to_response_middleware(
114 request: Request, call_next: Callable[[Request], Any]
115) -> Response:
116 trace_id = trace.get_current_span().get_span_context().trace_id
117 response = await call_next(request)
118 response.headers["Chroma-Trace-Id"] = format(trace_id, "x")
119 return response
120
121
122async def catch_exceptions_middleware(
123 request: Request, call_next: Callable[[Request], Any]
124) -> Response:
125 try:
126 return await call_next(request)
127 except ChromaError as e:
128 return fastapi_json_response(e)
129 except ValueError as e:
130 return ORJSONResponse(
131 content={"error": "InvalidArgumentError", "message": str(e)},
132 status_code=400,
133 )
134 except TypeError as e:
135 return ORJSONResponse(
136 content={"error": "InvalidArgumentError", "message": str(e)},
137 status_code=400,
138 )
139 except Exception as e:
140 logger.exception(e)
141 return ORJSONResponse(content={"error": repr(e)}, status_code=500)
142
143
144async def check_http_version_middleware(
145 request: Request, call_next: Callable[[Request], Any]
146) -> Response:
147 http_version = request.scope.get("http_version")
148 if http_version not in ["1.1", "2"]:
149 raise InvalidHTTPVersion(f"HTTP version {http_version} is not supported")
150 return await call_next(request)
151
152
153D = TypeVar("D", bound=BaseModel, contravariant=True)
154
155
156def validate_model(model: Type[D], data: Any) -> D: # type: ignore
157 """Used for backward compatibility with Pydantic 1.x"""
158 try:
159 return model.model_validate(data) # pydantic 2.x
160 except AttributeError:
161 return model.parse_obj(data) # pydantic 1.x
162
163
164class ChromaAPIRouter(fastapi.APIRouter): # type: ignore
165 # A simple subclass of fastapi's APIRouter which treats URLs with a
166 # trailing "/" the same as URLs without. Docs will only contain URLs
167 # without trailing "/"s.
168 def add_api_route(self, path: str, *args: Any, **kwargs: Any) -> None:
169 # If kwargs["include_in_schema"] isn't passed OR is True, we should
170 # only include the non-"/" path. If kwargs["include_in_schema"] is
171 # False, include neither.
172 exclude_from_schema = (
173 "include_in_schema" in kwargs and not kwargs["include_in_schema"]
174 )
175
176 def include_in_schema(path: str) -> bool:
177 nonlocal exclude_from_schema
178 return not exclude_from_schema and not path.endswith("/")
179
180 kwargs["include_in_schema"] = include_in_schema(path)
181 super().add_api_route(path, *args, **kwargs)
182
183 if path.endswith("/"):
184 path = path[:-1]
185 else:
186 path = path + "/"
187
188 kwargs["include_in_schema"] = include_in_schema(path)
189 super().add_api_route(path, *args, **kwargs)
190
191
192class FastAPI(Server):
193 def __init__(self, settings: Settings):
194 ProductTelemetryClient.SERVER_CONTEXT = ServerContext.FASTAPI
195 # https://fastapi.tiangolo.com/advanced/custom-response/#use-orjsonresponse
196 self._app = fastapi.FastAPI(debug=True, default_response_class=ORJSONResponse)
197 self._system = System(settings)
198 self._api: ServerAPI = self._system.instance(ServerAPI)
199
200 self._extra_openapi_schemas: Dict[str, Any] = {}
201 self._app.openapi = self.generate_openapi
202
203 self._opentelemetry_client = self._api.require(OpenTelemetryClient)
204 self._capacity_limiter = CapacityLimiter(
205 settings.chroma_server_thread_pool_size
206 )
207 self._quota_enforcer = self._system.require(QuotaEnforcer)
208 self._system.start()
209
210 self._app.middleware("http")(check_http_version_middleware)
211 self._app.middleware("http")(catch_exceptions_middleware)
212 self._app.middleware("http")(add_trace_id_to_response_middleware)
213 self._app.add_middleware(
214 CORSMiddleware,
215 allow_headers=["*"],
216 allow_origins=settings.chroma_server_cors_allow_origins,
217 allow_methods=["*"],
218 )
219 self._app.add_exception_handler(QuotaError, self.quota_exception_handler)
220 self._app.add_exception_handler(
221 RateLimitError, self.rate_limit_exception_handler
222 )
223 self._async_rate_limit_enforcer = self._system.require(AsyncRateLimitEnforcer)
224
225 self._app.on_event("shutdown")(self.shutdown)
226
227 self.authn_provider = None
228 if settings.chroma_server_authn_provider:
229 self.authn_provider = self._system.require(ServerAuthenticationProvider)
230
231 self.authz_provider = None
232 if settings.chroma_server_authz_provider:
233 self.authz_provider = self._system.require(ServerAuthorizationProvider)
234
235 self.router = ChromaAPIRouter()
236
237 self.setup_v1_routes()
238 self.setup_v2_routes()
239
240 self._app.include_router(self.router)
241
242 use_route_names_as_operation_ids(self._app)
243 instrument_fastapi(self._app)
244 telemetry_client = self._system.instance(ProductTelemetryClient)
245 telemetry_client.capture(ServerStartEvent())
246
247 def generate_openapi(self) -> Dict[str, Any]:
248 """Used instead of the default openapi() generation handler to include manually-populated schemas."""
249 schema: Dict[str, Any] = get_openapi(
250 title="Chroma",
251 routes=self._app.routes,
252 version=chromadb_version,
253 )
254
255 for key, value in self._extra_openapi_schemas.items():
256 schema["components"]["schemas"][key] = value
257
258 return schema
259
260 def get_openapi_extras_for_body_model(
261 self, request_model: Type[D]
262 ) -> Dict[str, Any]:
263 schema = request_model.model_json_schema(
264 ref_template="#/components/schemas/{model}"
265 )
266 if "$defs" in schema:
267 for key, value in schema["$defs"].items():
268 self._extra_openapi_schemas[key] = value
269
270 openapi_extra = {
271 "requestBody": {
272 "content": {"application/json": {"schema": schema}},
273 "required": True,
274 }
275 }
276 return openapi_extra
277
278 def setup_v2_routes(self) -> None:
279 self.router.add_api_route("/api/v2", self.root, methods=["GET"])
280 self.router.add_api_route("/api/v2/reset", self.reset, methods=["POST"])
281 self.router.add_api_route("/api/v2/version", self.version, methods=["GET"])
282 self.router.add_api_route("/api/v2/heartbeat", self.heartbeat, methods=["GET"])
283 self.router.add_api_route(
284 "/api/v2/pre-flight-checks", self.pre_flight_checks, methods=["GET"]
285 )
286
287 self.router.add_api_route(
288 "/api/v2/auth/identity",
289 self.get_user_identity,
290 methods=["GET"],
291 response_model=None,
292 )
293
294 self.router.add_api_route(
295 "/api/v2/tenants/{tenant}/databases",
296 self.create_database,
297 methods=["POST"],
298 response_model=None,
299 openapi_extra=self.get_openapi_extras_for_body_model(CreateDatabase),
300 )
301
302 self.router.add_api_route(
303 "/api/v2/tenants/{tenant}/databases/{database_name}",
304 self.get_database,
305 methods=["GET"],
306 response_model=None,
307 )
308
309 self.router.add_api_route(
310 "/api/v2/tenants/{tenant}/databases/{database_name}",
311 self.delete_database,
312 methods=["DELETE"],
313 response_model=None,
314 )
315
316 self.router.add_api_route(
317 "/api/v2/tenants",
318 self.create_tenant,
319 methods=["POST"],
320 response_model=None,
321 openapi_extra=self.get_openapi_extras_for_body_model(CreateTenant),
322 )
323
324 self.router.add_api_route(
325 "/api/v2/tenants/{tenant}",
326 self.get_tenant,
327 methods=["GET"],
328 response_model=None,
329 )
330
331 self.router.add_api_route(
332 "/api/v2/tenants/{tenant}/databases",
333 self.list_databases,
334 methods=["GET"],
335 response_model=None,
336 )
337
338 self.router.add_api_route(
339 "/api/v2/tenants/{tenant}/databases/{database_name}/collections",
340 self.list_collections,
341 methods=["GET"],
342 response_model=None,
343 )
344 self.router.add_api_route(
345 "/api/v2/tenants/{tenant}/databases/{database_name}/collections_count",
346 self.count_collections,
347 methods=["GET"],
348 response_model=None,
349 )
350 self.router.add_api_route(
351 "/api/v2/tenants/{tenant}/databases/{database_name}/collections",
352 self.create_collection,
353 methods=["POST"],
354 response_model=None,
355 openapi_extra=self.get_openapi_extras_for_body_model(CreateCollection),
356 )
357
358 self.router.add_api_route(
359 "/api/v2/tenants/{tenant}/databases/{database_name}/collections/{collection_id}/add",
360 self.add,
361 methods=["POST"],
362 status_code=status.HTTP_201_CREATED,
363 response_model=None,
364 openapi_extra=self.get_openapi_extras_for_body_model(AddEmbedding),
365 )
366 self.router.add_api_route(
367 "/api/v2/tenants/{tenant}/databases/{database_name}/collections/{collection_id}/update",
368 self.update,
369 methods=["POST"],
370 response_model=None,
371 openapi_extra=self.get_openapi_extras_for_body_model(UpdateEmbedding),
372 )
373 self.router.add_api_route(
374 "/api/v2/tenants/{tenant}/databases/{database_name}/collections/{collection_id}/upsert",
375 self.upsert,
376 methods=["POST"],
377 response_model=None,
378 openapi_extra=self.get_openapi_extras_for_body_model(AddEmbedding),
379 )
380 self.router.add_api_route(
381 "/api/v2/tenants/{tenant}/databases/{database_name}/collections/{collection_id}/get",
382 self.get,
383 methods=["POST"],
384 response_model=None,
385 openapi_extra=self.get_openapi_extras_for_body_model(GetEmbedding),
386 )
387 self.router.add_api_route(
388 "/api/v2/tenants/{tenant}/databases/{database_name}/collections/{collection_id}/delete",
389 self.delete,
390 methods=["POST"],
391 response_model=DeleteResult,
392 openapi_extra=self.get_openapi_extras_for_body_model(DeleteEmbedding),
393 )
394 self.router.add_api_route(
395 "/api/v2/tenants/{tenant}/databases/{database_name}/collections/{collection_id}/count",
396 self.count,
397 methods=["GET"],
398 response_model=None,
399 )
400 self.router.add_api_route(
401 "/api/v2/tenants/{tenant}/databases/{database_name}/collections/{collection_id}/query",
402 self.get_nearest_neighbors,
403 methods=["POST"],
404 response_model=None,
405 openapi_extra=self.get_openapi_extras_for_body_model(
406 request_model=QueryEmbedding
407 ),
408 )
409 self.router.add_api_route(
410 "/api/v2/tenants/{tenant}/databases/{database_name}/collections/{collection_name}",
411 self.get_collection,
412 methods=["GET"],
413 response_model=None,
414 )
415 self.router.add_api_route(
416 "/api/v2/tenants/{tenant}/databases/{database_name}/collections/{collection_id}",
417 self.update_collection,
418 methods=["PUT"],
419 response_model=None,
420 openapi_extra=self.get_openapi_extras_for_body_model(UpdateCollection),
421 )
422 self.router.add_api_route(
423 "/api/v2/tenants/{tenant}/databases/{database_name}/collections/{collection_name}",
424 self.delete_collection,
425 methods=["DELETE"],
426 response_model=None,
427 )
428
429 self.router.add_api_route(
430 "/api/v2/tenants/{tenant}/databases/{database_name}/collections/{collection_id}/functions/attach",
431 self.attach_function,
432 methods=["POST"],
433 response_model=None,
434 )
435
436 self.router.add_api_route(
437 "/api/v2/tenants/{tenant}/databases/{database_name}/collections/{collection_id}/functions/{function_name}",
438 self.get_attached_function,
439 methods=["GET"],
440 response_model=None,
441 )
442
443 self.router.add_api_route(
444 "/api/v2/tenants/{tenant}/databases/{database_name}/collections/by-id/{collection_id}",
445 self.get_collection_by_id,
446 methods=["GET"],
447 response_model=None,
448 )
449
450 def shutdown(self) -> None:
451 self._system.stop()
452
453 def app(self) -> fastapi.FastAPI:
454 return self._app
455
456 async def rate_limit_exception_handler(
457 self, request: Request, exc: RateLimitError
458 ) -> ORJSONResponse:
459 return ORJSONResponse(
460 status_code=429,
461 content={"message": "Rate limit exceeded."},
462 )
463
464 def root(self) -> Dict[str, int]:
465 return {"nanosecond heartbeat": self._api.heartbeat()}
466
467 async def quota_exception_handler(
468 self, request: Request, exc: QuotaError
469 ) -> ORJSONResponse:
470 return ORJSONResponse(
471 status_code=400,
472 content={"message": exc.message()},
473 )
474
475 async def heartbeat(self) -> Dict[str, int]:
476 return self.root()
477
478 async def version(self) -> str:
479 return self._api.get_version()
480
481 def _set_request_context(self, request: Request) -> None:
482 """
483 Set context about the request on any components that might need it.
484 """
485 self._quota_enforcer.set_context(context={"request": request})
486
487 @trace_method(
488 "auth_request",
489 OpenTelemetryGranularity.OPERATION,
490 )
491 @rate_limit
492 async def auth_request(
493 self,
494 headers: Headers,
495 action: AuthzAction,
496 tenant: Optional[str],
497 database: Optional[str],
498 collection: Optional[str],
499 ) -> None:
500 return await to_thread.run_sync(
501 # NOTE(rescrv, iron will auth): No need to migrate because this is the utility call.
502 self.sync_auth_request,
503 *(headers, action, tenant, database, collection),
504 )
505
506 @trace_method(
507 "FastAPI.sync_auth_request",
508 OpenTelemetryGranularity.OPERATION,
509 )
510 def sync_auth_request(
511 self,
512 headers: Headers,
513 action: AuthzAction,
514 tenant: Optional[str],
515 database: Optional[str],
516 collection: Optional[str],
517 ) -> None:
518 """
519 Authenticates and authorizes the request based on the given headers
520 and other parameters. If the request cannot be authenticated or cannot
521 be authorized (with the configured providers), raises an HTTP 401.
522 """
523 if not self.authn_provider:
524 add_attributes_to_current_span(
525 {
526 "tenant": tenant,
527 "database": database,
528 "collection": collection,
529 }
530 )
531 return
532
533 user_identity = self.authn_provider.authenticate_or_raise(dict(headers))
534
535 if not self.authz_provider:
536 return
537
538 authz_resource = AuthzResource(
539 tenant=tenant,
540 database=database,
541 collection=collection,
542 )
543
544 self.authz_provider.authorize_or_raise(user_identity, action, authz_resource)
545 add_attributes_to_current_span(
546 {
547 "tenant": tenant,
548 "database": database,
549 "collection": collection,
550 }
551 )
552 return
553
554 @trace_method("FastAPI.get_user_identity", OpenTelemetryGranularity.OPERATION)
555 async def get_user_identity(
556 self,
557 request: Request,
558 ) -> UserIdentity:
559 if not self.authn_provider:
560 return UserIdentity(
561 user_id="", tenant=DEFAULT_TENANT, databases=[DEFAULT_DATABASE]
562 )
563
564 return cast(
565 UserIdentity,
566 await to_thread.run_sync(
567 lambda: cast(
568 ServerAuthenticationProvider, self.authn_provider
569 ).authenticate_or_raise(dict(request.headers)) # type: ignore
570 ),
571 )
572
573 @trace_method("FastAPI.create_database", OpenTelemetryGranularity.OPERATION)
574 async def create_database(
575 self,
576 request: Request,
577 tenant: str,
578 ) -> None:
579 def process_create_database(
580 tenant: str, headers: Headers, raw_body: bytes
581 ) -> None:
582 db = validate_model(CreateDatabase, orjson.loads(raw_body))
583
584 # NOTE(rescrv, iron will auth): Implemented.
585 self.sync_auth_request(
586 headers,
587 AuthzAction.CREATE_DATABASE,
588 tenant,
589 db.name,
590 None,
591 )
592
593 self._set_request_context(request=request)
594
595 return self._api.create_database(db.name, tenant)
596
597 await to_thread.run_sync(
598 process_create_database,
599 tenant,
600 request.headers,
601 await request.body(),
602 limiter=self._capacity_limiter,
603 )
604
605 @trace_method("FastAPI.get_database", OpenTelemetryGranularity.OPERATION)
606 async def get_database(
607 self,
608 request: Request,
609 database_name: str,
610 tenant: str,
611 ) -> Database:
612 # NOTE(rescrv, iron will auth): Implemented.
613 await self.auth_request(
614 request.headers,
615 AuthzAction.GET_DATABASE,
616 tenant,
617 database_name,
618 None,
619 )
620
621 return cast(
622 Database,
623 await to_thread.run_sync(
624 self._api.get_database,
625 database_name,
626 tenant,
627 limiter=self._capacity_limiter,
628 ),
629 )
630
631 @trace_method("FastAPI.delete_database", OpenTelemetryGranularity.OPERATION)
632 async def delete_database(
633 self,
634 request: Request,
635 database_name: str,
636 tenant: str,
637 ) -> None:
638 # NOTE(rescrv, iron will auth): Implemented.
639 await self.auth_request(
640 request.headers,
641 AuthzAction.DELETE_DATABASE,
642 tenant,
643 database_name,
644 None,
645 )
646
647 await to_thread.run_sync(
648 self._api.delete_database,
649 database_name,
650 tenant,
651 limiter=self._capacity_limiter,
652 )
653
654 @trace_method("FastAPI.create_tenant", OpenTelemetryGranularity.OPERATION)
655 async def create_tenant(
656 self,
657 request: Request,
658 ) -> None:
659 def process_create_tenant(request: Request, raw_body: bytes) -> None:
660 tenant = validate_model(CreateTenant, orjson.loads(raw_body))
661
662 # NOTE(rescrv, iron will auth): Implemented.
663 self.sync_auth_request(
664 request.headers,
665 AuthzAction.CREATE_TENANT,
666 tenant.name,
667 None,
668 None,
669 )
670
671 return self._api.create_tenant(tenant.name)
672
673 await to_thread.run_sync(
674 process_create_tenant,
675 request,
676 await request.body(),
677 limiter=self._capacity_limiter,
678 )
679
680 @trace_method("FastAPI.get_tenant", OpenTelemetryGranularity.OPERATION)
681 async def get_tenant(
682 self,
683 request: Request,
684 tenant: str,
685 ) -> Tenant:
686 # NOTE(rescrv, iron will auth): Implemented.
687 await self.auth_request(
688 request.headers,
689 AuthzAction.GET_TENANT,
690 tenant,
691 None,
692 None,
693 )
694
695 return cast(
696 Tenant,
697 await to_thread.run_sync(
698 self._api.get_tenant,
699 tenant,
700 limiter=self._capacity_limiter,
701 ),
702 )
703
704 @trace_method("FastAPI.list_databases", OpenTelemetryGranularity.OPERATION)
705 async def list_databases(
706 self,
707 request: Request,
708 tenant: str,
709 limit: Optional[int] = None,
710 offset: Optional[int] = None,
711 ) -> Sequence[Database]:
712 # NOTE(rescrv, iron will auth): Implemented.
713 await self.auth_request(
714 request.headers,
715 AuthzAction.LIST_DATABASES,
716 tenant,
717 None,
718 None,
719 )
720
721 return cast(
722 Sequence[Database],
723 await to_thread.run_sync(
724 self._api.list_databases,
725 limit,
726 offset,
727 tenant,
728 limiter=self._capacity_limiter,
729 ),
730 )
731
732 @trace_method("FastAPI.list_collections", OpenTelemetryGranularity.OPERATION)
733 async def list_collections(
734 self,
735 request: Request,
736 tenant: str,
737 database_name: str,
738 limit: Optional[int] = None,
739 offset: Optional[int] = None,
740 ) -> Sequence[CollectionModel]:
741 def process_list_collections(
742 limit: Optional[int], offset: Optional[int], tenant: str, database_name: str
743 ) -> Sequence[CollectionModel]:
744 # NOTE(rescrv, iron will auth): Implemented.
745 self.sync_auth_request(
746 request.headers,
747 AuthzAction.LIST_COLLECTIONS,
748 tenant,
749 database_name,
750 None,
751 )
752
753 self._set_request_context(request=request)
754
755 add_attributes_to_current_span({"tenant": tenant})
756 return self._api.list_collections(
757 tenant=tenant, database=database_name, limit=limit, offset=offset
758 )
759
760 api_collection_models = cast(
761 Sequence[CollectionModel],
762 await to_thread.run_sync(
763 process_list_collections,
764 limit,
765 offset,
766 tenant,
767 database_name,
768 limiter=self._capacity_limiter,
769 ),
770 )
771
772 return api_collection_models
773
774 @trace_method("FastAPI.count_collections", OpenTelemetryGranularity.OPERATION)
775 async def count_collections(
776 self,
777 request: Request,
778 tenant: str,
779 database_name: str,
780 ) -> int:
781 # NOTE(rescrv, iron will auth): Implemented.
782 await self.auth_request(
783 request.headers,
784 AuthzAction.COUNT_COLLECTIONS,
785 tenant,
786 database_name,
787 None,
788 )
789
790 add_attributes_to_current_span({"tenant": tenant})
791
792 return cast(
793 int,
794 await to_thread.run_sync(
795 self._api.count_collections,
796 tenant,
797 database_name,
798 limiter=self._capacity_limiter,
799 ),
800 )
801
802 @trace_method("FastAPI.create_collection", OpenTelemetryGranularity.OPERATION)
803 async def create_collection(
804 self,
805 request: Request,
806 tenant: str,
807 database_name: str,
808 ) -> CollectionModel:
809 def process_create_collection(
810 request: Request, tenant: str, database: str, raw_body: bytes
811 ) -> CollectionModel:
812 create = validate_model(CreateCollection, orjson.loads(raw_body))
813 if not create.configuration:
814 if create.metadata:
815 configuration = (
816 create_collection_configuration_from_legacy_collection_metadata(
817 create.metadata
818 )
819 )
820 else:
821 configuration = None
822 else:
823 configuration = load_create_collection_configuration_from_json(
824 create.configuration
825 )
826
827 # NOTE(rescrv, iron will auth): Implemented.
828 self.sync_auth_request(
829 request.headers,
830 AuthzAction.CREATE_COLLECTION,
831 tenant,
832 database,
833 create.name,
834 )
835
836 self._set_request_context(request=request)
837
838 add_attributes_to_current_span({"tenant": tenant})
839
840 return self._api.create_collection(
841 name=create.name,
842 configuration=configuration,
843 metadata=create.metadata,
844 get_or_create=create.get_or_create,
845 tenant=tenant,
846 database=database,
847 )
848
849 api_collection_model = cast(
850 CollectionModel,
851 await to_thread.run_sync(
852 process_create_collection,
853 request,
854 tenant,
855 database_name,
856 await request.body(),
857 limiter=self._capacity_limiter,
858 ),
859 )
860 return api_collection_model
861
862 @trace_method("FastAPI.get_collection", OpenTelemetryGranularity.OPERATION)
863 async def get_collection(
864 self,
865 request: Request,
866 tenant: str,
867 database_name: str,
868 collection_name: str,
869 ) -> CollectionModel:
870 # NOTE(rescrv, iron will auth): Implemented.
871 await self.auth_request(
872 request.headers,
873 AuthzAction.GET_COLLECTION,
874 tenant,
875 database_name,
876 collection_name,
877 )
878
879 add_attributes_to_current_span({"tenant": tenant})
880
881 api_collection_model = cast(
882 CollectionModel,
883 await to_thread.run_sync(
884 self._api.get_collection,
885 collection_name,
886 tenant,
887 database_name,
888 limiter=self._capacity_limiter,
889 ),
890 )
891 return api_collection_model
892
893 @trace_method("FastAPI.get_collection_by_id", OpenTelemetryGranularity.OPERATION)
894 async def get_collection_by_id(
895 self,
896 request: Request,
897 tenant: str,
898 database_name: str,
899 collection_id: str,
900 ) -> CollectionModel:
901 # NOTE(rescrv, iron will auth): Implemented.
902 await self.auth_request(
903 request.headers,
904 AuthzAction.GET_COLLECTION,
905 tenant,
906 database_name,
907 collection_id,
908 )
909
910 add_attributes_to_current_span({"tenant": tenant})
911
912 api_collection_model = cast(
913 CollectionModel,
914 await to_thread.run_sync(
915 self._api.get_collection_by_id,
916 _uuid(collection_id),
917 tenant,
918 database_name,
919 limiter=self._capacity_limiter,
920 ),
921 )
922 return api_collection_model
923
924 @trace_method("FastAPI.update_collection", OpenTelemetryGranularity.OPERATION)
925 async def update_collection(
926 self,
927 tenant: str,
928 database_name: str,
929 collection_id: str,
930 request: Request,
931 ) -> None:
932 def process_update_collection(
933 request: Request, collection_id: str, raw_body: bytes
934 ) -> None:
935 update = validate_model(UpdateCollection, orjson.loads(raw_body))
936 # NOTE(rescrv, iron will auth): Implemented.
937 self.sync_auth_request(
938 request.headers,
939 AuthzAction.UPDATE_COLLECTION,
940 tenant,
941 database_name,
942 collection_id,
943 )
944 configuration = (
945 None
946 if not update.new_configuration
947 else load_update_collection_configuration_from_json(
948 update.new_configuration
949 )
950 )
951 self._set_request_context(request=request)
952 add_attributes_to_current_span({"tenant": tenant})
953 return self._api._modify(
954 id=_uuid(collection_id),
955 new_name=update.new_name,
956 new_metadata=update.new_metadata,
957 new_configuration=configuration,
958 tenant=tenant,
959 database=database_name,
960 )
961
962 await to_thread.run_sync(
963 process_update_collection,
964 request,
965 collection_id,
966 await request.body(),
967 limiter=self._capacity_limiter,
968 )
969
970 @trace_method("FastAPI.delete_collection", OpenTelemetryGranularity.OPERATION)
971 async def delete_collection(
972 self,
973 request: Request,
974 collection_name: str,
975 tenant: str,
976 database_name: str,
977 ) -> None:
978 # NOTE(rescrv, iron will auth): Implemented.
979 await self.auth_request(
980 request.headers,
981 AuthzAction.DELETE_COLLECTION,
982 tenant,
983 database_name,
984 collection_name,
985 )
986 add_attributes_to_current_span({"tenant": tenant})
987
988 await to_thread.run_sync(
989 self._api.delete_collection,
990 collection_name,
991 tenant,
992 database_name,
993 limiter=self._capacity_limiter,
994 )
995
996 @trace_method("FastAPI.attach_function", OpenTelemetryGranularity.OPERATION)
997 @rate_limit
998 async def attach_function(
999 self,
1000 request: Request,
1001 tenant: str,
1002 database_name: str,
1003 collection_id: str,
1004 ) -> Dict[str, Any]:
1005 try:
1006
1007 def process_attach_function(
1008 request: Request, raw_body: bytes
1009 ) -> Dict[str, Any]:
1010 body = orjson.loads(raw_body)
1011 # NOTE: Auth check for attaching functions
1012 self.sync_auth_request(
1013 request.headers,
1014 AuthzAction.UPDATE_COLLECTION, # Using UPDATE_COLLECTION as the auth action
1015 tenant,
1016 database_name,
1017 collection_id,
1018 )
1019 self._set_request_context(request=request)
1020
1021 name = body.get("name")
1022 function_id = body.get("function_id")
1023 output_collection = body.get("output_collection")
1024 params = body.get("params")
1025
1026 attached_fn = self._api.attach_function(
1027 function_id=function_id,
1028 name=name,
1029 input_collection_id=_uuid(collection_id),
1030 output_collection=output_collection,
1031 params=params,
1032 tenant=tenant,
1033 database=database_name,
1034 )
1035
1036 return {
1037 "attached_function": {
1038 "id": str(attached_fn.id),
1039 "name": attached_fn.name,
1040 "function_name": attached_fn.function_name,
1041 "output_collection": attached_fn.output_collection,
1042 "params": attached_fn.params,
1043 }
1044 }
1045
1046 raw_body = await request.body()
1047 return await to_thread.run_sync(
1048 process_attach_function,
1049 request,
1050 raw_body,
1051 limiter=self._capacity_limiter,
1052 )
1053 except Exception:
1054 raise
1055
1056 @trace_method("FastAPI.get_attached_function", OpenTelemetryGranularity.OPERATION)
1057 @rate_limit
1058 async def get_attached_function(
1059 self,
1060 request: Request,
1061 tenant: str,
1062 database_name: str,
1063 collection_id: str,
1064 function_name: str,
1065 ) -> Dict[str, Any]:
1066 # NOTE: Auth check for getting attached functions
1067 await self.auth_request(
1068 request.headers,
1069 AuthzAction.GET_COLLECTION, # Using GET_COLLECTION as the auth action
1070 tenant,
1071 database_name,
1072 collection_id,
1073 )
1074 add_attributes_to_current_span({"tenant": tenant})
1075
1076 attached_fn = await to_thread.run_sync(
1077 self._api.get_attached_function,
1078 function_name,
1079 _uuid(collection_id),
1080 tenant,
1081 database_name,
1082 limiter=self._capacity_limiter,
1083 )
1084
1085 return {
1086 "attached_function": {
1087 "id": str(attached_fn.id),
1088 "name": attached_fn.name,
1089 "function_name": attached_fn.function_name,
1090 "output_collection": attached_fn.output_collection,
1091 "params": attached_fn.params,
1092 }
1093 }
1094
1095 @trace_method("FastAPI.add", OpenTelemetryGranularity.OPERATION)
1096 @rate_limit
1097 async def add(
1098 self,
1099 request: Request,
1100 tenant: str,
1101 database_name: str,
1102 collection_id: str,
1103 ) -> bool:
1104 try:
1105
1106 def process_add(request: Request, raw_body: bytes) -> bool:
1107 add = validate_model(AddEmbedding, orjson.loads(raw_body))
1108 # NOTE(rescrv, iron will auth): Implemented.
1109 self.sync_auth_request(
1110 request.headers,
1111 AuthzAction.ADD,
1112 tenant,
1113 database_name,
1114 collection_id,
1115 )
1116 self._set_request_context(request=request)
1117 add_attributes_to_current_span({"tenant": tenant})
1118 return self._api._add(
1119 collection_id=_uuid(collection_id),
1120 ids=add.ids,
1121 embeddings=cast(
1122 Embeddings,
1123 convert_list_embeddings_to_np(add.embeddings)
1124 if add.embeddings
1125 else None,
1126 ),
1127 metadatas=add.metadatas, # type: ignore
1128 documents=add.documents, # type: ignore
1129 uris=add.uris, # type: ignore
1130 tenant=tenant,
1131 database=database_name,
1132 )
1133
1134 return cast(
1135 bool,
1136 await to_thread.run_sync(
1137 process_add,
1138 request,
1139 await request.body(),
1140 limiter=self._capacity_limiter,
1141 ),
1142 )
1143 except InvalidDimensionException as e:
1144 raise HTTPException(status_code=500, detail=str(e))
1145
1146 @trace_method("FastAPI.update", OpenTelemetryGranularity.OPERATION)
1147 @rate_limit
1148 async def update(
1149 self,
1150 request: Request,
1151 tenant: str,
1152 database_name: str,
1153 collection_id: str,
1154 ) -> None:
1155 def process_update(request: Request, raw_body: bytes) -> bool:
1156 update = validate_model(UpdateEmbedding, orjson.loads(raw_body))
1157
1158 # NOTE(rescrv, iron will auth): Implemented.
1159 self.sync_auth_request(
1160 request.headers,
1161 AuthzAction.UPDATE,
1162 tenant,
1163 database_name,
1164 collection_id,
1165 )
1166 self._set_request_context(request=request)
1167 add_attributes_to_current_span({"tenant": tenant})
1168
1169 return self._api._update(
1170 collection_id=_uuid(collection_id),
1171 ids=update.ids,
1172 embeddings=convert_list_embeddings_to_np(update.embeddings)
1173 if update.embeddings
1174 else None,
1175 metadatas=update.metadatas, # type: ignore
1176 documents=update.documents, # type: ignore
1177 uris=update.uris, # type: ignore
1178 tenant=tenant,
1179 database=database_name,
1180 )
1181
1182 await to_thread.run_sync(
1183 process_update,
1184 request,
1185 await request.body(),
1186 limiter=self._capacity_limiter,
1187 )
1188
1189 @trace_method("FastAPI.upsert", OpenTelemetryGranularity.OPERATION)
1190 @rate_limit
1191 async def upsert(
1192 self,
1193 request: Request,
1194 tenant: str,
1195 database_name: str,
1196 collection_id: str,
1197 ) -> None:
1198 def process_upsert(request: Request, raw_body: bytes) -> bool:
1199 upsert = validate_model(AddEmbedding, orjson.loads(raw_body))
1200
