|
| 1 | +import json as json_lib |
| 2 | +from typing import Type |
| 3 | + |
| 4 | +import google.genai as genai |
| 5 | +from google.genai import types as genai_types |
| 6 | +from pydantic import ValidationError |
| 7 | +from typing_extensions import override |
| 8 | + |
| 9 | +from askui.logger import logger |
| 10 | +from askui.models.askui.inference_api import AskUiInferenceApiSettings |
| 11 | +from askui.models.exceptions import QueryNoResponseError, QueryUnexpectedResponseError |
| 12 | +from askui.models.models import GetModel, ModelName |
| 13 | +from askui.models.shared.prompts import SYSTEM_PROMPT_GET |
| 14 | +from askui.models.types.response_schemas import ResponseSchema, to_response_schema |
| 15 | +from askui.utils.image_utils import ImageSource |
| 16 | + |
| 17 | +ASKUI_MODEL_CHOICE_PREFIX = "askui/" |
| 18 | +ASKUI_MODEL_CHOICE_PREFIX_LEN = len(ASKUI_MODEL_CHOICE_PREFIX) |
| 19 | + |
| 20 | + |
| 21 | +def _extract_model_id(model_choice: str) -> str: |
| 22 | + if model_choice == ModelName.ASKUI: |
| 23 | + return ModelName.GEMINI__2_5__FLASH |
| 24 | + if model_choice.startswith(ASKUI_MODEL_CHOICE_PREFIX): |
| 25 | + return model_choice[ASKUI_MODEL_CHOICE_PREFIX_LEN:] |
| 26 | + return model_choice |
| 27 | + |
| 28 | + |
| 29 | +class AskUiGoogleGenAiApi(GetModel): |
| 30 | + def __init__(self, settings: AskUiInferenceApiSettings | None = None) -> None: |
| 31 | + self._settings = settings or AskUiInferenceApiSettings() |
| 32 | + self._client = genai.Client( |
| 33 | + vertexai=True, |
| 34 | + api_key="Necessary", |
| 35 | + http_options=genai_types.HttpOptions( |
| 36 | + base_url=f"{self._settings.base_url}/proxy/vertexai", |
| 37 | + headers={ |
| 38 | + "Authorization": self._settings.authorization_header, |
| 39 | + }, |
| 40 | + ), |
| 41 | + ) |
| 42 | + |
| 43 | + @override |
| 44 | + def get( |
| 45 | + self, |
| 46 | + query: str, |
| 47 | + image: ImageSource, |
| 48 | + response_schema: Type[ResponseSchema] | None, |
| 49 | + model_choice: str, |
| 50 | + ) -> ResponseSchema | str: |
| 51 | + try: |
| 52 | + _response_schema = to_response_schema(response_schema) |
| 53 | + json_schema = _response_schema.model_json_schema() |
| 54 | + logger.debug(f"json_schema:\n{json_lib.dumps(json_schema)}") |
| 55 | + content = genai_types.Content( |
| 56 | + parts=[ |
| 57 | + genai_types.Part.from_bytes( |
| 58 | + data=image.to_bytes(), |
| 59 | + mime_type="image/png", |
| 60 | + ), |
| 61 | + genai_types.Part.from_text(text=query), |
| 62 | + ], |
| 63 | + role="user", |
| 64 | + ) |
| 65 | + generate_content_response = self._client.models.generate_content( |
| 66 | + model=f"models/{_extract_model_id(model_choice)}", |
| 67 | + contents=content, |
| 68 | + config={ |
| 69 | + "response_mime_type": "application/json", |
| 70 | + "response_schema": _response_schema, |
| 71 | + "system_instruction": SYSTEM_PROMPT_GET, |
| 72 | + }, |
| 73 | + ) |
| 74 | + json_str = generate_content_response.text |
| 75 | + if json_str is None: |
| 76 | + raise QueryNoResponseError( |
| 77 | + message="No response from the model", query=query |
| 78 | + ) |
| 79 | + try: |
| 80 | + return _response_schema.model_validate_json(json_str).root |
| 81 | + except ValidationError as e: |
| 82 | + error_message = str(e.errors()) |
| 83 | + raise QueryUnexpectedResponseError( |
| 84 | + message=f"Unexpected response from the model: {error_message}", |
| 85 | + query=query, |
| 86 | + response=json_str, |
| 87 | + ) from e |
| 88 | + except RecursionError as e: |
| 89 | + error_message = ( |
| 90 | + "Recursive response schemas are not supported by AskUiGoogleGenAiApi" |
| 91 | + ) |
| 92 | + raise NotImplementedError(error_message) from e |
0 commit comments