Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
__init__.py2335 linesDownload Raw Back to fastapi
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

Showing the first 1,200 of 2335 lines. Download the file for the rest.

codekingpro/portable-devtools · Team Ai