first shot

This commit is contained in:
9qeklajc
2025-11-22 20:07:26 +01:00
parent 14ae4ecce3
commit 0f23335de9
7 changed files with 3005 additions and 1490 deletions
+1
View File
@@ -20,6 +20,7 @@ dependencies = [
"nostr>=0.0.2", "nostr>=0.0.2",
"mdurl==0.1.2", "mdurl==0.1.2",
"pillow>=10", "pillow>=10",
"google-generativeai>=0.8.5",
] ]
[dependency-groups] [dependency-groups]
+2
View File
@@ -2,6 +2,7 @@ from .anthropic import AnthropicUpstreamProvider
from .azure import AzureUpstreamProvider from .azure import AzureUpstreamProvider
from .base import BaseUpstreamProvider from .base import BaseUpstreamProvider
from .fireworks import FireworksUpstreamProvider from .fireworks import FireworksUpstreamProvider
from .gemini import GeminiUpstreamProvider
from .generic import GenericUpstreamProvider from .generic import GenericUpstreamProvider
from .groq import GroqUpstreamProvider from .groq import GroqUpstreamProvider
from .ollama import OllamaUpstreamProvider from .ollama import OllamaUpstreamProvider
@@ -14,6 +15,7 @@ upstream_provider_classes: list[type[BaseUpstreamProvider]] = [
AnthropicUpstreamProvider, AnthropicUpstreamProvider,
AzureUpstreamProvider, AzureUpstreamProvider,
FireworksUpstreamProvider, FireworksUpstreamProvider,
GeminiUpstreamProvider,
GenericUpstreamProvider, GenericUpstreamProvider,
GroqUpstreamProvider, GroqUpstreamProvider,
OllamaUpstreamProvider, OllamaUpstreamProvider,
+3
View File
@@ -0,0 +1,3 @@
from .gemini import GeminiClient
__all__ = ["GeminiClient"]
+41
View File
@@ -0,0 +1,41 @@
from __future__ import annotations
from abc import ABC, abstractmethod
from typing import Any, AsyncGenerator
class BaseAPIClient(ABC):
"""Base class for AI provider API clients."""
def __init__(self, api_key: str, base_url: str | None = None):
self.api_key = api_key
self.base_url = base_url
@abstractmethod
async def generate_content(
self,
model: str,
messages: list[dict[str, Any]],
temperature: float | None = None,
max_tokens: int | None = None,
**kwargs: Any,
) -> dict[str, Any]:
"""Generate content non-streaming."""
pass
@abstractmethod
async def generate_content_stream(
self,
model: str,
messages: list[dict[str, Any]],
temperature: float | None = None,
max_tokens: int | None = None,
**kwargs: Any,
) -> AsyncGenerator[dict[str, Any], None]:
"""Generate content with streaming."""
pass
@abstractmethod
async def list_models(self) -> list[dict[str, Any]]:
"""List available models."""
pass
+202
View File
@@ -0,0 +1,202 @@
from __future__ import annotations
import time
from typing import Any, AsyncGenerator
import google.generativeai as genai
from .base import BaseAPIClient
class GeminiClient(BaseAPIClient):
"""Native Gemini API client using Google's official package."""
def __init__(self, api_key: str, base_url: str | None = None):
super().__init__(api_key, base_url)
genai.configure(api_key=api_key) # type: ignore
def _validate_model_name(self, model: str) -> str:
return model
def _convert_openai_to_gemini_messages(self, messages: list[dict[str, Any]]) -> str:
"""Convert OpenAI messages to a simple string for Gemini."""
combined_content = []
for message in messages:
role = message.get("role", "user")
content = message.get("content", "")
if role == "system":
combined_content.append(f"System: {content}")
elif role == "user":
combined_content.append(f"User: {content}")
elif role == "assistant":
combined_content.append(f"Assistant: {content}")
return "\n".join(combined_content)
async def generate_content(
self,
model: str,
messages: list[dict[str, Any]],
temperature: float | None = None,
max_tokens: int | None = None,
**kwargs: Any,
) -> dict[str, Any]:
"""Generate content using Gemini API (non-streaming)."""
model = self._validate_model_name(model)
prompt = self._convert_openai_to_gemini_messages(messages)
generation_config = {}
if temperature is not None:
generation_config["temperature"] = temperature
if max_tokens is not None:
generation_config["max_output_tokens"] = max_tokens
if "top_p" in kwargs:
generation_config["top_p"] = kwargs["top_p"]
model_instance = genai.GenerativeModel(model) # type: ignore
response = await model_instance.generate_content_async( # type: ignore
prompt,
generation_config=generation_config if generation_config else None, # type: ignore
)
return {
"id": f"chatcmpl-{hash(str(response))}"[1:16],
"object": "chat.completion",
"created": int(time.time()),
"model": model,
"choices": [
{
"index": 0,
"message": {
"role": "assistant",
"content": response.text,
},
"finish_reason": "stop",
}
],
"usage": {
"prompt_tokens": getattr(
response.usage_metadata, "prompt_token_count", 0
),
"completion_tokens": getattr(
response.usage_metadata, "candidates_token_count", 0
),
"total_tokens": getattr(
response.usage_metadata, "total_token_count", 0
),
},
}
async def generate_content_stream( # type: ignore[override]
self,
model: str,
messages: list[dict[str, Any]],
temperature: float | None = None,
max_tokens: int | None = None,
**kwargs: Any,
) -> AsyncGenerator[dict[str, Any], None]:
"""Generate content using Gemini API (streaming)."""
model = self._validate_model_name(model)
prompt = self._convert_openai_to_gemini_messages(messages)
stream_id = f"chatcmpl-{abs(hash(prompt + str(time.time())))}"[:28]
created_time = int(time.time())
generation_config = {}
if temperature is not None:
generation_config["temperature"] = temperature
if max_tokens is not None:
generation_config["max_output_tokens"] = max_tokens
if "top_p" in kwargs:
generation_config["top_p"] = kwargs["top_p"]
model_instance = genai.GenerativeModel(model) # type: ignore
response_stream = await model_instance.generate_content_async( # type: ignore
prompt,
generation_config=generation_config if generation_config else None, # type: ignore
stream=True,
)
async for chunk in response_stream:
finish_reason = None
content = ""
if hasattr(chunk, "text") and chunk.text:
content = chunk.text
if hasattr(chunk, "candidates") and chunk.candidates:
candidate = chunk.candidates[0]
finish_reason_raw = getattr(candidate, "finish_reason", None)
if finish_reason_raw == 1:
finish_reason = "stop"
elif finish_reason_raw == 2:
finish_reason = "length"
elif finish_reason_raw == 3:
finish_reason = "content_filter"
elif finish_reason_raw == 4:
finish_reason = "content_filter"
usage_data = None
if hasattr(chunk, "usage_metadata") and chunk.usage_metadata:
usage_metadata = chunk.usage_metadata
usage_data = {
"prompt_tokens": getattr(usage_metadata, "prompt_token_count", 0),
"completion_tokens": getattr(
usage_metadata, "candidates_token_count", 0
),
"total_tokens": getattr(usage_metadata, "total_token_count", 0),
}
if hasattr(usage_metadata, "cached_content_token_count"):
cached_tokens = getattr(
usage_metadata, "cached_content_token_count", 0
)
if cached_tokens > 0:
usage_data["prompt_tokens_details"] = {
"cached_tokens": cached_tokens
}
chunk_data = {
"id": stream_id,
"object": "chat.completion.chunk",
"created": created_time,
"model": model,
"choices": [
{
"index": 0,
"delta": {"content": content} if content else {},
"finish_reason": finish_reason,
}
],
}
if usage_data:
chunk_data["usage"] = usage_data
yield chunk_data
if finish_reason == "stop":
break
async def list_models(self) -> list[dict[str, Any]]:
"""List available Gemini models."""
try:
models = genai.list_models() # type: ignore
return [
{
"name": model.name,
"displayName": getattr(model, "display_name", model.name),
"description": getattr(model, "description", ""),
"supportedGenerationMethods": getattr(
model, "supported_generation_methods", ["generateContent"]
),
"inputTokenLimit": getattr(model, "input_token_limit", 32768),
"outputTokenLimit": getattr(model, "output_token_limit", 8192),
}
for model in models
]
except Exception as e:
from ...core.logging import get_logger
logger = get_logger(__name__)
logger.error(f"Failed to list Gemini models: {e}")
return []
+583
View File
@@ -0,0 +1,583 @@
from __future__ import annotations
import json
from collections.abc import AsyncGenerator
from typing import TYPE_CHECKING
from fastapi import Request
from fastapi.responses import Response, StreamingResponse
from .base import BaseUpstreamProvider
from .clients.gemini import GeminiClient
if TYPE_CHECKING:
from ..core.db import ApiKey, AsyncSession, UpstreamProviderRow
from ..payment.models import Model
from ..core.logging import get_logger
logger = get_logger(__name__)
class GeminiUpstreamProvider(BaseUpstreamProvider):
provider_type = "gemini"
default_base_url = "https://generativelanguage.googleapis.com/v1beta"
platform_url = "https://aistudio.google.com/app/apikey"
def __init__(
self,
base_url: str = "https://generativelanguage.googleapis.com/v1beta",
api_key: str = "",
provider_fee: float = 1.01,
):
super().__init__(
api_key=api_key,
provider_fee=provider_fee,
base_url=base_url,
)
self._client: GeminiClient | None = None
@property
def client(self) -> GeminiClient:
"""Get or create the Gemini API client."""
if self._client is None:
self._client = GeminiClient(api_key=self.api_key)
return self._client
@classmethod
def from_db_row(
cls, provider_row: "UpstreamProviderRow"
) -> "GeminiUpstreamProvider":
return cls(
base_url=provider_row.base_url,
api_key=provider_row.api_key,
provider_fee=provider_row.provider_fee,
)
@classmethod
def get_provider_metadata(cls) -> dict[str, object]:
return {
"id": cls.provider_type,
"name": "Google Gemini",
"default_base_url": cls.default_base_url,
"fixed_base_url": True,
"platform_url": cls.platform_url,
}
def prepare_headers(self, request_headers: dict) -> dict:
headers = dict(request_headers)
removed_headers = []
for header in [
"host",
"content-length",
"refund-lnurl",
"key-expiry-time",
"x-cashu",
"authorization",
]:
if headers.pop(header, None) is not None:
removed_headers.append(header)
if self.api_key:
headers["x-goog-api-key"] = self.api_key
return headers
def transform_model_name(self, model_id: str) -> str:
return model_id.removeprefix("gemini/")
def prepare_request_body(
self, body: bytes | None, model_obj: Model
) -> bytes | None:
if not body:
return body
try:
openai_data = json.loads(body)
if not isinstance(openai_data, dict):
return body
gemini_data = {}
if "messages" in openai_data:
contents = []
for message in openai_data["messages"]:
role = message.get("role", "user")
content = message.get("content", "")
if role == "system":
continue
elif role == "assistant":
role = "model"
contents.append({"role": role, "parts": [{"text": content}]})
gemini_data["contents"] = contents
generation_config = {}
if "temperature" in openai_data:
generation_config["temperature"] = openai_data["temperature"]
if "max_tokens" in openai_data:
generation_config["maxOutputTokens"] = openai_data["max_tokens"]
if "top_p" in openai_data:
generation_config["topP"] = openai_data["top_p"]
if generation_config:
gemini_data["generationConfig"] = generation_config # type: ignore
return json.dumps(gemini_data).encode()
except (json.JSONDecodeError, KeyError, TypeError) as e:
logger.warning(
f"Failed to transform request body for Gemini: {e}",
extra={"error": str(e), "error_type": type(e).__name__},
)
return body
async def forward_request(
self,
request: Request,
path: str,
headers: dict,
request_body: bytes | None,
key: ApiKey,
max_cost_for_model: int,
session: AsyncSession,
model_obj: Model,
) -> Response | StreamingResponse:
if not path.startswith("chat/completions"):
return await super().forward_request(
request,
path,
headers,
request_body,
key,
max_cost_for_model,
session,
model_obj,
)
if not request_body:
return await super().forward_request(
request,
path,
headers,
request_body,
key,
max_cost_for_model,
session,
model_obj,
)
try:
openai_data = json.loads(request_body)
messages = openai_data.get("messages", [])
temperature = openai_data.get("temperature")
max_tokens = openai_data.get("max_tokens")
top_p = openai_data.get("top_p")
is_streaming = openai_data.get("stream", False)
logger.info(
"Processing Gemini request with client abstraction",
extra={
"model": model_obj.id,
"is_streaming": is_streaming,
"message_count": len(messages),
"key_hash": key.hashed_key[:8] + "...",
},
)
if is_streaming:
response_generator = self.client.generate_content_stream(
model=model_obj.id,
messages=messages,
temperature=temperature,
max_tokens=max_tokens,
top_p=top_p,
)
async def stream_with_cost() -> AsyncGenerator[bytes, None]:
last_model_seen: str = model_obj.id
final_usage_data: dict | None = None
try:
async for chunk in response_generator:
if chunk.get("usage"):
final_usage_data = chunk["usage"]
sse_data = f"data: {json.dumps(chunk)}\n\n"
yield sse_data.encode()
if (
chunk.get("choices", [{}])[0].get("finish_reason")
== "stop"
):
break
payment_data = {
"model": last_model_seen,
"usage": final_usage_data,
}
from ..auth import adjust_payment_for_tokens
from ..core.db import create_session
async with create_session() as new_session:
fresh_key = await new_session.get(
key.__class__, key.hashed_key
)
if fresh_key:
try:
cost_data = await adjust_payment_for_tokens(
fresh_key,
payment_data,
new_session,
max_cost_for_model,
)
logger.info(
"Gemini streaming payment finalized",
extra={
"cost_data": cost_data,
"usage_data": final_usage_data,
"key_hash": key.hashed_key[:8] + "...",
},
)
except Exception as cost_error:
logger.error(
"Error finalizing Gemini streaming payment",
extra={
"error": str(cost_error),
"key_hash": key.hashed_key[:8] + "...",
},
)
except Exception as e:
logger.error(
"Error in Gemini streaming response",
extra={
"error": str(e),
"error_type": type(e).__name__,
"key_hash": key.hashed_key[:8] + "...",
},
)
raise
return StreamingResponse(
stream_with_cost(),
media_type="text/event-stream",
headers={"Cache-Control": "no-cache", "Connection": "keep-alive"},
)
else:
openai_format_response = await self.client.generate_content(
model=model_obj.id,
messages=messages,
temperature=temperature,
max_tokens=max_tokens,
top_p=top_p,
)
from ..auth import adjust_payment_for_tokens
cost_data = await adjust_payment_for_tokens(
key, openai_format_response, session, max_cost_for_model
)
openai_format_response["cost"] = cost_data
logger.info(
"Gemini non-streaming payment completed",
extra={
"cost_data": cost_data,
"model": model_obj.id,
"key_hash": key.hashed_key[:8] + "...",
},
)
return Response(
content=json.dumps(openai_format_response),
media_type="application/json",
headers={"Cache-Control": "no-cache"},
)
except Exception as e:
logger.error(
"Error in Gemini forward_request",
extra={
"error": str(e),
"error_type": type(e).__name__,
"path": path,
"key_hash": key.hashed_key[:8] + "...",
},
)
return await super().forward_request(
request,
path,
headers,
request_body,
key,
max_cost_for_model,
session,
model_obj,
)
def _transform_gemini_to_openai(self, gemini_response: dict, model_id: str) -> dict:
"""Transform Gemini API response to OpenAI format."""
candidates = gemini_response.get("candidates", [])
if not candidates:
return {
"id": f"chatcmpl-{hash(str(gemini_response))}"[1:16],
"object": "chat.completion",
"created": int(__import__("time").time()),
"model": model_id,
"choices": [],
"usage": {
"prompt_tokens": 0,
"completion_tokens": 0,
"total_tokens": 0,
},
}
first_candidate = candidates[0]
content = ""
if "content" in first_candidate and "parts" in first_candidate["content"]:
parts = first_candidate["content"]["parts"]
if parts and "text" in parts[0]:
content = parts[0]["text"]
finish_reason = first_candidate.get("finishReason", "stop")
if finish_reason == "STOP":
finish_reason = "stop"
elif finish_reason == "MAX_TOKENS":
finish_reason = "length"
usage_metadata = gemini_response.get("usageMetadata", {})
prompt_tokens = usage_metadata.get("promptTokenCount", 0)
completion_tokens = usage_metadata.get("candidatesTokenCount", 0)
return {
"id": f"chatcmpl-{hash(str(gemini_response))}"[1:16],
"object": "chat.completion",
"created": int(__import__("time").time()),
"model": model_id,
"choices": [
{
"index": 0,
"message": {
"role": "assistant",
"content": content,
},
"finish_reason": finish_reason,
}
],
"usage": {
"prompt_tokens": prompt_tokens,
"completion_tokens": completion_tokens,
"total_tokens": prompt_tokens + completion_tokens,
},
}
async def fetch_models(self) -> list[Model]:
from ..payment.models import Architecture, Model, Pricing, TopProvider
try:
models_data = await self.client.list_models()
models_list = []
for model_data in models_data:
model_name = model_data.get("name", "")
if not model_name:
continue
model_id = model_name.replace("models/", "")
display_name = model_data.get("displayName", model_id)
description = model_data.get(
"description", f"Google {display_name} model"
)
supported_methods = model_data.get("supportedGenerationMethods", [])
if "generateContent" not in supported_methods:
logger.debug(
f"Skipping model {model_id} - doesn't support generateContent",
extra={"supported_methods": supported_methods},
)
continue
input_token_limit = model_data.get("inputTokenLimit", 32768)
output_token_limit = model_data.get("outputTokenLimit", 8192)
context_length = min(input_token_limit, 128000)
logger.debug(
f"Found Gemini model: {model_id}",
extra={
"display_name": display_name,
"supported_methods": supported_methods,
"input_token_limit": input_token_limit,
"output_token_limit": output_token_limit,
"context_length": context_length,
},
)
pricing_config = Pricing(
prompt=0.000003,
completion=0.000003,
request=0.0,
image=0.0,
web_search=0.0,
internal_reasoning=0.0,
max_prompt_cost=0.001,
max_completion_cost=0.001,
max_cost=0.001,
)
if "1.5-pro" in model_id.lower() or "2.0" in model_id.lower():
pricing_config = Pricing(
prompt=0.00000125,
completion=0.000005,
request=0.0,
image=0.0,
web_search=0.0,
internal_reasoning=0.0,
max_prompt_cost=0.001,
max_completion_cost=0.001,
max_cost=0.001,
)
elif (
"1.5-flash" in model_id.lower()
or "2.5" in model_id.lower()
or "flash" in model_id.lower()
):
pricing_config = Pricing(
prompt=0.000000075,
completion=0.0000003,
request=0.0,
image=0.0,
web_search=0.0,
internal_reasoning=0.0,
max_prompt_cost=0.001,
max_completion_cost=0.001,
max_cost=0.001,
)
models_list.append(
Model(
id=model_id,
name=display_name,
created=0,
description=description,
context_length=context_length,
architecture=Architecture(
modality="text",
input_modalities=["text"],
output_modalities=["text"],
tokenizer="gemini",
instruct_type=None,
),
pricing=pricing_config,
sats_pricing=None,
per_request_limits=None,
top_provider=TopProvider(
context_length=context_length,
max_completion_tokens=output_token_limit,
is_moderated=True,
),
enabled=True,
upstream_provider_id=None,
canonical_slug=None,
)
)
logger.info(
f"Fetched {len(models_list)} models from Gemini",
extra={"model_count": len(models_list), "base_url": self.base_url},
)
return models_list
except Exception as e:
logger.error(
f"Failed to fetch models from Gemini API: {e}",
extra={
"error": str(e),
"error_type": type(e).__name__,
"base_url": self.base_url,
},
)
return []
async def refresh_models_cache(self) -> None:
try:
from ..payment.models import _update_model_sats_pricing
from ..payment.price import sats_usd_price
models = await self.fetch_models()
models_with_fees = [self._apply_provider_fee_to_model(m) for m in models]
try:
sats_to_usd = sats_usd_price()
self._models_cache = [
_update_model_sats_pricing(m, sats_to_usd) for m in models_with_fees
]
except Exception:
self._models_cache = models_with_fees
self._models_by_id = {m.id: m for m in self._models_cache}
logger.info(
f"Refreshed models cache for {self.base_url}",
extra={"model_count": len(models)},
)
except Exception as e:
logger.error(
f"Failed to refresh models cache for {self.base_url}",
extra={"error": str(e), "error_type": type(e).__name__},
)
def get_cached_models(self) -> list[Model]:
return self._models_cache
def get_cached_model_by_id(self, model_id: str) -> Model | None:
return self._models_by_id.get(model_id)
def _apply_provider_fee_to_model(self, model: Model) -> Model:
from ..payment.models import Model, Pricing, _calculate_usd_max_costs
adjusted_pricing = Pricing.parse_obj(
{k: v * self.provider_fee for k, v in model.pricing.dict().items()}
)
temp_model = Model(
id=model.id,
name=model.name,
created=model.created,
description=model.description,
context_length=model.context_length,
architecture=model.architecture,
pricing=adjusted_pricing,
sats_pricing=None,
per_request_limits=model.per_request_limits,
top_provider=model.top_provider,
enabled=model.enabled,
upstream_provider_id=model.upstream_provider_id,
canonical_slug=model.canonical_slug,
)
(
adjusted_pricing.max_prompt_cost,
adjusted_pricing.max_completion_cost,
adjusted_pricing.max_cost,
) = _calculate_usd_max_costs(temp_model)
return Model(
id=model.id,
name=model.name,
created=model.created,
description=model.description,
context_length=model.context_length,
architecture=model.architecture,
pricing=adjusted_pricing,
sats_pricing=model.sats_pricing,
per_request_limits=model.per_request_limits,
top_provider=model.top_provider,
enabled=model.enabled,
upstream_provider_id=model.upstream_provider_id,
canonical_slug=model.canonical_slug,
)
Generated
+2173 -1490
View File
File diff suppressed because it is too large Load Diff