Spaces:
Runtime error
Runtime error
| """ | |
| Gemini API Pool Manager | |
| Alternative to Groq when Groq is blocked in user's region | |
| Uses Gemini 2.x models for agent reasoning | |
| """ | |
| import asyncio | |
| import time | |
| import logging | |
| from collections import deque | |
| from typing import List, Dict | |
| from dataclasses import dataclass | |
| logger = logging.getLogger(__name__) | |
| class GeminiKeyUsage: | |
| """Track usage for a Gemini API key""" | |
| api_key: str | |
| calls_today: int = 0 | |
| last_reset_date: str = "" | |
| total_calls: int = 0 | |
| total_failures: int = 0 | |
| is_blocked: bool = False | |
| blocked_until: float = 0 | |
| class GeminiAgentPool: | |
| """ | |
| Gemini-based agent pool (alternative to Groq) | |
| Uses Gemini 2.x models for fast agent reasoning | |
| Free tier: 1500 requests/day for gemini-2.0-flash | |
| """ | |
| # Available models - PRIORITIZE HIGHER RPM MODELS FOR FREE TIER | |
| # Free tier limits (per key): | |
| # gemini-2.5-flash: 10 RPM | |
| # gemini-2.0-flash: 15 RPM | |
| # gemini-2.5-flash-lite: 15 RPM | |
| # gemini-2.0-flash-lite: 30 RPM <-- BEST for high concurrency | |
| REASONING_MODEL = "gemini-2.0-flash" # 15 RPM, good reasoning | |
| FAST_MODEL = "gemini-2.0-flash-lite" # 30 RPM, fast routing | |
| LITE_MODEL = "gemini-2.0-flash-lite" # 30 RPM | |
| # Fallback model chain (highest RPM first) | |
| MODEL_FALLBACKS = [ | |
| "gemini-2.0-flash-lite", # 30 RPM | |
| "gemini-2.5-flash-lite", # 15 RPM | |
| "gemini-2.0-flash", # 15 RPM | |
| "gemini-2.5-flash", # 10 RPM | |
| "gemini-flash-latest", | |
| ] | |
| def __init__(self, api_keys: List[str]): | |
| """Initialize with Gemini API keys""" | |
| if not api_keys: | |
| raise ValueError("At least one Gemini API key is required") | |
| self.api_keys = api_keys | |
| self.key_usage = {key: GeminiKeyUsage(api_key=key) for key in api_keys} | |
| self.current_index = 0 | |
| self._lock = asyncio.Lock() | |
| self._configured_keys = set() | |
| self._working_model = None | |
| # Per-key rate limit tracking (timestamps of recent calls) | |
| self._key_call_times = {key: deque() for key in api_keys} | |
| # Concurrency limiter - avoid bursting too many concurrent calls | |
| # With 3 keys × 30 RPM = 90 calls/min capacity, allow 6 concurrent | |
| max_concurrent = min(6, len(api_keys) * 2) | |
| self._semaphore = asyncio.Semaphore(max_concurrent) | |
| logger.info( | |
| f"GeminiAgentPool initialized: {len(api_keys)} keys, " | |
| f"max concurrent={max_concurrent}" | |
| ) | |
| def _configure_key(self, api_key: str): | |
| """Configure Gemini SDK with this API key""" | |
| if api_key not in self._configured_keys: | |
| import google.generativeai as genai | |
| genai.configure(api_key=api_key) | |
| self._configured_keys.add(api_key) | |
| async def get_available_key(self) -> str: | |
| """Get next available API key (round-robin)""" | |
| async with self._lock: | |
| attempts = 0 | |
| max_attempts = len(self.api_keys) * 2 | |
| while attempts < max_attempts: | |
| key = self.api_keys[self.current_index % len(self.api_keys)] | |
| self.current_index += 1 | |
| usage = self.key_usage[key] | |
| # Reset block if expired | |
| now = time.time() | |
| if usage.is_blocked and now > usage.blocked_until: | |
| usage.is_blocked = False | |
| if not usage.is_blocked: | |
| usage.total_calls += 1 | |
| return key | |
| attempts += 1 | |
| await asyncio.sleep(2) | |
| return await self.get_available_key() | |
| async def invoke( | |
| self, | |
| prompt: str, | |
| model: str = "llama-3.1-8b-instant", # Ignored, mapped to Gemini | |
| temperature: float = 0.2, | |
| max_tokens: int = 1000, | |
| system_prompt: str = None, | |
| agent_name: str = "unknown", | |
| conversation_id: str = None, | |
| ) -> str: | |
| """ | |
| Invoke Gemini API (compatible with GroqAPIPool interface) | |
| """ | |
| from .llm_logger import llm_logger | |
| # Map Groq model names to Gemini models | |
| gemini_model = self._map_model(model) | |
| api_key = await self.get_available_key() | |
| start_time = time.time() | |
| try: | |
| import google.generativeai as genai | |
| self._configure_key(api_key) | |
| generation_config = { | |
| "temperature": temperature, | |
| "max_output_tokens": max_tokens, | |
| } | |
| models_to_try = [self._working_model] if self._working_model else [] | |
| models_to_try.extend([m for m in self.MODEL_FALLBACKS if m != self._working_model]) | |
| models_to_try = [m for m in models_to_try if m] | |
| last_error = None | |
| for model_name in models_to_try: | |
| try: | |
| gen_model = genai.GenerativeModel( | |
| model_name=model_name, | |
| system_instruction=system_prompt if system_prompt else None, | |
| generation_config=generation_config, | |
| ) | |
| response = await asyncio.to_thread( | |
| gen_model.generate_content, | |
| prompt, | |
| ) | |
| self._working_model = model_name | |
| response_text = response.text if hasattr(response, "text") else str(response) | |
| duration_ms = (time.time() - start_time) * 1000 | |
| # Log this call | |
| llm_logger.log_call( | |
| agent_name=agent_name, | |
| provider="gemini", | |
| model=model_name, | |
| prompt=prompt, | |
| response=response_text, | |
| tokens_input=len(prompt) // 4, | |
| tokens_output=len(response_text) // 4, | |
| duration_ms=duration_ms, | |
| success=True, | |
| conversation_id=conversation_id, | |
| metadata={"system_prompt_chars": len(system_prompt) if system_prompt else 0}, | |
| ) | |
| return response_text | |
| except Exception as e: | |
| last_error = e | |
| err_str = str(e).lower() | |
| # If model not found, try next | |
| if "404" in err_str or "not found" in err_str: | |
| continue | |
| # If rate limited, mark key as blocked and retry with next key | |
| if "429" in err_str or "rate" in err_str or "quota" in err_str: | |
| usage = self.key_usage[api_key] | |
| usage.is_blocked = True | |
| # Parse retry_delay from error if possible | |
| retry_seconds = 30 | |
| try: | |
| if "retry_delay" in err_str: | |
| import re | |
| m = re.search(r"seconds:\s*(\d+)", err_str) | |
| if m: | |
| retry_seconds = int(m.group(1)) | |
| except Exception: | |
| pass | |
| usage.blocked_until = time.time() + retry_seconds | |
| logger.warning( | |
| f"Gemini key {api_key[:10]}... rate limited (retry in {retry_seconds}s). " | |
| f"Trying next key from pool ({len(self.api_keys)} total)." | |
| ) | |
| # Try next key (recursion will pick a different key) | |
| return await self.invoke( | |
| prompt, model, temperature, max_tokens, system_prompt, | |
| agent_name, conversation_id | |
| ) | |
| # Other errors - try next model | |
| continue | |
| raise Exception(f"All Gemini models failed. Last error: {last_error}") | |
| except Exception as e: | |
| usage = self.key_usage[api_key] | |
| usage.total_failures += 1 | |
| logger.error(f"Gemini invocation failed: {e}") | |
| raise | |
| def _map_model(self, groq_model: str) -> str: | |
| """Map Groq-style model name to Gemini equivalent""" | |
| # 70B = needs reasoning -> use gemini-2.5-flash | |
| # 8B = fast routing -> use gemini-2.0-flash | |
| if "70b" in groq_model.lower() or "versatile" in groq_model.lower(): | |
| return self.REASONING_MODEL | |
| elif "8b" in groq_model.lower() or "instant" in groq_model.lower(): | |
| return self.FAST_MODEL | |
| return self.FAST_MODEL # Default | |
| def get_status(self) -> Dict: | |
| """Get current pool status""" | |
| return { | |
| "provider": "gemini", | |
| "total_keys": len(self.api_keys), | |
| "active_keys": sum(1 for k in self.key_usage.values() if not k.is_blocked), | |
| "total_calls": sum(k.total_calls for k in self.key_usage.values()), | |
| "total_failures": sum(k.total_failures for k in self.key_usage.values()), | |
| "working_model": self._working_model, | |
| } | |