from __future__ import annotations from typing import Any, Mapping import requests from .base import ( AiCapabilityError, AiResponseParseError, ImageGenerationResult, ResolvedModel, TextGenerationResult, validate_model_config, ) from .utils import ( API_CHAT, API_GEMINI, API_IMAGES, API_IMAGES_EDITS, extract_image_from_response, extract_text_from_response, extract_titles_from_response, image_bytes_to_data_url, normalize_api_url, request_timeout, resolution_to_size, split_data_url, ) SAFE_PARAMETER_KEYS = frozenset( { "temperature", "top_p", "topP", "top_k", "topK", "max_tokens", "max_output_tokens", "maxOutputTokens", "presence_penalty", "presencePenalty", "frequency_penalty", "frequencyPenalty", "seed", "stop", } ) SAFE_GENERATION_CONFIG_KEYS = frozenset( { "temperature", "topP", "topK", "maxOutputTokens", "stopSequences", } ) class BaseHttpProvider: def __init__(self, session: requests.Session | None = None): self.session = session or requests.Session() if hasattr(self.session, "trust_env"): self.session.trust_env = False def _headers(self, model: ResolvedModel, *, json: bool = False) -> dict[str, str]: headers = {"Authorization": f"Bearer {model.api_key}"} if json: headers["Content-Type"] = "application/json" return headers def _timeout(self, model: ResolvedModel, resolution: str) -> tuple[int, int]: return request_timeout( model.connect_timeout_seconds, model.timeout_seconds, resolution, ) def _read_timeout(self, model: ResolvedModel, resolution: str) -> int: return self._timeout(model, resolution)[1] class ChatCompletionsProvider(BaseHttpProvider): def capabilities(self) -> set[str]: return {"text", "image", "vision"} def generate_text( self, prompt: str, model: ResolvedModel, *, image: bytes | None = None, image_mime_type: str = "image/png", resolution: str = "1K", parameters: Mapping[str, Any] | None = None, ) -> TextGenerationResult: validate_model_config(model) url = normalize_api_url(model.url, API_CHAT) payload = build_chat_text_payload( model, prompt, image=image, image_mime_type=image_mime_type, parameters=parameters, ) response = self.session.post( url, headers=self._headers(model, json=True), json=payload, timeout=self._timeout(model, resolution), ) response.raise_for_status() raw = response.json() titles = extract_titles_from_response(raw) text = extract_text_from_response(raw) if not text: raise AiResponseParseError("AI response did not contain text") return TextGenerationResult(text=text, titles=titles, model_used=model.model, raw=raw) def generate_image( self, prompt: str, model: ResolvedModel, *, image: bytes | None = None, image_mime_type: str = "image/png", image_filename: str = "image.png", resolution: str = "1K", aspect_ratio: str = "1:1", parameters: Mapping[str, Any] | None = None, ) -> ImageGenerationResult: validate_model_config(model) url = normalize_api_url(model.url, API_CHAT) payload = build_chat_image_payload( model, prompt, image=image, image_mime_type=image_mime_type, parameters=parameters, ) response = self.session.post( url, headers=self._headers(model, json=True), json=payload, timeout=self._timeout(model, resolution), ) response.raise_for_status() raw = response.json() image_bytes = extract_image_from_response( raw, session=self.session, timeout=self._read_timeout(model, resolution), ) if not image_bytes: raise AiResponseParseError("AI response did not contain an image") return ImageGenerationResult(image=image_bytes, model_used=model.model, raw=raw) class GeminiProvider(ChatCompletionsProvider): def generate_text( self, prompt: str, model: ResolvedModel, *, image: bytes | None = None, image_mime_type: str = "image/png", resolution: str = "1K", parameters: Mapping[str, Any] | None = None, ) -> TextGenerationResult: validate_model_config(model) url = normalize_api_url(model.url, API_GEMINI).replace("{model}", model.model) payload = build_gemini_payload( model, prompt, image=image, image_mime_type=image_mime_type, response_modalities=["TEXT"], parameters=parameters, ) response = self.session.post( url, headers=self._headers(model, json=True), json=payload, timeout=self._timeout(model, resolution), ) response.raise_for_status() raw = response.json() titles = extract_titles_from_response(raw) text = extract_text_from_response(raw) if not text: raise AiResponseParseError("AI response did not contain text") return TextGenerationResult(text=text, titles=titles, model_used=model.model, raw=raw) def generate_image( self, prompt: str, model: ResolvedModel, *, image: bytes | None = None, image_mime_type: str = "image/png", image_filename: str = "image.png", resolution: str = "1K", aspect_ratio: str = "1:1", parameters: Mapping[str, Any] | None = None, ) -> ImageGenerationResult: validate_model_config(model) url = normalize_api_url(model.url, API_GEMINI).replace("{model}", model.model) payload = build_gemini_payload( model, prompt, image=image, image_mime_type=image_mime_type, response_modalities=["TEXT", "IMAGE"], parameters=parameters, ) response = self.session.post( url, headers=self._headers(model, json=True), json=payload, timeout=self._timeout(model, resolution), ) response.raise_for_status() raw = response.json() image_bytes = extract_image_from_response( raw, session=self.session, timeout=self._read_timeout(model, resolution), ) if not image_bytes: raise AiResponseParseError("AI response did not contain an image") return ImageGenerationResult(image=image_bytes, model_used=model.model, raw=raw) class ImagesGenerationProvider(BaseHttpProvider): def capabilities(self) -> set[str]: return {"image"} def generate_text(self, *args: Any, **kwargs: Any) -> TextGenerationResult: raise AiCapabilityError("images generation provider cannot generate text") def generate_image( self, prompt: str, model: ResolvedModel, *, image: bytes | None = None, image_mime_type: str = "image/png", image_filename: str = "image.png", resolution: str = "1K", aspect_ratio: str = "1:1", parameters: Mapping[str, Any] | None = None, ) -> ImageGenerationResult: validate_model_config(model) url = normalize_api_url(model.url, API_IMAGES) payload: dict[str, Any] = { "model": model.model, "prompt": prompt, "aspect_ratio": aspect_ratio, "resolution": resolution, "n": 1, } if image is not None: payload["image_urls"] = [image_bytes_to_data_url(image, image_mime_type)] apply_extra_body(payload, model, parameters) response = self.session.post( url, headers=self._headers(model, json=True), json=payload, timeout=self._timeout(model, resolution), ) response.raise_for_status() raw = response.json() image_bytes = extract_image_from_response( raw, session=self.session, timeout=self._read_timeout(model, resolution), ) if not image_bytes: raise AiResponseParseError("AI response did not contain an image") return ImageGenerationResult(image=image_bytes, model_used=model.model, raw=raw) class ImagesEditsProvider(BaseHttpProvider): def capabilities(self) -> set[str]: return {"image", "vision"} def generate_text(self, *args: Any, **kwargs: Any) -> TextGenerationResult: raise AiCapabilityError("images edits provider cannot generate text") def generate_image( self, prompt: str, model: ResolvedModel, *, image: bytes | None = None, image_mime_type: str = "image/png", image_filename: str = "image.png", resolution: str = "1K", aspect_ratio: str = "1:1", parameters: Mapping[str, Any] | None = None, ) -> ImageGenerationResult: validate_model_config(model) if image is None: raise AiCapabilityError("images edits provider requires an input image") url = normalize_api_url(model.url, API_IMAGES_EDITS) data: dict[str, Any] = { "model": model.model, "prompt": prompt, "n": "1", "size": resolution_to_size(resolution), } apply_extra_body(data, model, parameters) files = {"image": (image_filename, image, image_mime_type)} response = self.session.post( url, headers=self._headers(model), data=data, files=files, timeout=self._timeout(model, resolution), ) response.raise_for_status() raw = response.json() image_bytes = extract_image_from_response( raw, session=self.session, timeout=self._read_timeout(model, resolution), ) if not image_bytes: raise AiResponseParseError("AI response did not contain an image") return ImageGenerationResult(image=image_bytes, model_used=model.model, raw=raw) def build_chat_text_payload( model: ResolvedModel, prompt: str, *, image: bytes | None = None, image_mime_type: str = "image/png", parameters: Mapping[str, Any] | None = None, ) -> dict[str, Any]: content: list[dict[str, Any]] = [{"type": "text", "text": prompt}] if image is not None: content.append( { "type": "image_url", "image_url": {"url": image_bytes_to_data_url(image, image_mime_type)}, } ) payload: dict[str, Any] = { "model": model.model, "messages": [{"role": "user", "content": content}], "stream": False, } apply_extra_body(payload, model, parameters) return payload def build_chat_image_payload( model: ResolvedModel, prompt: str, *, image: bytes | None = None, image_mime_type: str = "image/png", parameters: Mapping[str, Any] | None = None, ) -> dict[str, Any]: return build_chat_text_payload( model, prompt, image=image, image_mime_type=image_mime_type, parameters=parameters, ) def build_gemini_payload( model: ResolvedModel, prompt: str, *, image: bytes | None = None, image_mime_type: str = "image/png", response_modalities: list[str], parameters: Mapping[str, Any] | None = None, ) -> dict[str, Any]: parts: list[dict[str, Any]] = [{"text": prompt}] if image is not None: data_url = image_bytes_to_data_url(image, image_mime_type) mime_type, data = split_data_url(data_url) parts.append({"inlineData": {"mimeType": mime_type, "data": data}}) payload: dict[str, Any] = { "contents": [{"parts": parts}], "generationConfig": {"responseModalities": response_modalities}, } apply_extra_body(payload, model, parameters) return payload def apply_extra_body( payload: dict[str, Any], model: ResolvedModel, parameters: Mapping[str, Any] | None = None, ) -> None: apply_safe_parameters(payload, model.extra_body) if parameters: apply_safe_parameters(payload, parameters) def apply_safe_parameters(payload: dict[str, Any], values: Mapping[str, Any]) -> None: for key, value in values.items(): if key == "generationConfig" and isinstance(value, Mapping): generation_config = payload.setdefault("generationConfig", {}) if not isinstance(generation_config, dict): continue for config_key, config_value in value.items(): if config_key in SAFE_GENERATION_CONFIG_KEYS: generation_config[config_key] = config_value continue if key in SAFE_PARAMETER_KEYS: payload[key] = value