fund-flow-backend / src /copilot /gemini_pool.py
Aniket2006's picture
feat(copilot): complete copilot integration and resolve CSS styling compliance
f70ac6a
Raw
History Blame Contribute Delete
9.22 kB
"""
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__)
@dataclass
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,
}