CoolFace
Apppublic

alschameri/helping-source

sourceHugging Faceapache-2.0updated 1y agoView on Hugging Face
0likes
gemini_client.py213 linesDownload Raw Back to root
1"""2Gemini API client for embeddings and text generation.3"""4 5import os6import time7import logging8from typing import List, Dict, Optional9import numpy as np10import requests11import json12 13logger = logging.getLogger(__name__)14 15class GeminiClient:16    """Client for interacting with Google Gemini API."""17    18    def __init__(self):19        self.api_key = os.getenv('GEMINI_API_KEY')20        if not self.api_key:21            raise ValueError("GEMINI_API_KEY environment variable is required")22        23        self.project_id = os.getenv('GEMINI_PROJECT', '')24        self.base_url = "https://generativelanguage.googleapis.com/v1beta"25        26        # Rate limiting27        self.last_request_time = 028        self.min_request_interval = 1.0  # seconds29        30        logger.info("Gemini client initialized")31    32    def _wait_for_rate_limit(self):33        """Simple rate limiting to avoid hitting API limits."""34        current_time = time.time()35        time_since_last = current_time - self.last_request_time36        37        if time_since_last < self.min_request_interval:38            sleep_time = self.min_request_interval - time_since_last39            time.sleep(sleep_time)40        41        self.last_request_time = time.time()42    43    def _make_request(self, url: str, payload: Dict, retries: int = 3) -> Dict:44        """Make HTTP request to Gemini API with retry logic."""45        for attempt in range(retries):46            try:47                self._wait_for_rate_limit()48                49                headers = {50                    'Content-Type': 'application/json'51                }52                53                response = requests.post(54                    f"{url}?key={self.api_key}",55                    headers=headers,56                    json=payload,57                    timeout=3058                )59                60                if response.status_code == 200:61                    return response.json()62                elif response.status_code == 429:  # Rate limit63                    wait_time = (2 ** attempt) * 2  # Exponential backoff64                    logger.warning(f"Rate limited, waiting {wait_time}s before retry {attempt + 1}")65                    time.sleep(wait_time)66                    continue67                else:68                    logger.error(f"API request failed: {response.status_code} - {response.text}")69                    response.raise_for_status()70                    71            except requests.exceptions.RequestException as e:72                logger.error(f"Request attempt {attempt + 1} failed: {e}")73                if attempt == retries - 1:74                    raise75                time.sleep(2 ** attempt)76        77        raise Exception("All retry attempts failed")78    79    def embed_texts(self, texts: List[str]) -> List[np.ndarray]:80        """Generate embeddings for a list of texts using Gemini."""81        if not texts:82            return []83        84        try:85            # Gemini embedding endpoint86            url = f"{self.base_url}/models/text-embedding-004:embedContent"87            88            embeddings = []89            90            # Process texts in batches to avoid hitting limits91            batch_size = 1092            for i in range(0, len(texts), batch_size):93                batch_texts = texts[i:i + batch_size]94                95                for text in batch_texts:96                    payload = {97                        "model": "models/text-embedding-004",98                        "content": {99                            "parts": [{100                                "text": text101                            }]102                        }103                    }104                    105                    response_data = self._make_request(url, payload)106                    107                    if 'embedding' in response_data and 'values' in response_data['embedding']:108                        embedding = np.array(response_data['embedding']['values'], dtype=np.float32)109                        embeddings.append(embedding)110                    else:111                        logger.error(f"Unexpected embedding response: {response_data}")112                        # Fallback to random embedding113                        embeddings.append(np.random.rand(self.embedding_dim).astype(np.float32))114            115            logger.info(f"Generated {len(embeddings)} embeddings")116            return embeddings117            118        except Exception as e:119            logger.error(f"Error generating embeddings: {e}")120            # Fallback to random embeddings for development121            logger.warning("Using random embeddings as fallback")122            return [np.random.rand(self.embedding_dim).astype(np.float32) for _ in texts]123    124    def generate_with_context(self, system_prompt: str, user_message: str, contexts: List[str]) -> Dict:125        """Generate response using Gemini with provided context."""126        try:127            # Build the complete prompt128            context_section = ""129            if contexts:130                context_section = "\n\nالسياق المتاح:\n" + "\n---\n".join(contexts)131            132            full_prompt = f"""{system_prompt}133 134{context_section}135 136سؤال المستخدم: {user_message}137 138يرجى الإجابة باللغة العربية مع اقتراح 2-4 أسئلة متابعة مفيدة. اجعل إجابتك مفيدة ومختصرة (2-4 جمل)."""139 140            # Gemini generation endpoint141            url = f"{self.base_url}/models/gemini-1.5-flash:generateContent"142            143            payload = {144                "contents": [{145                    "parts": [{146                        "text": full_prompt147                    }]148                }],149                "generationConfig": {150                    "temperature": 0.7,151                    "topK": 40,152                    "topP": 0.95,153                    "maxOutputTokens": 512,154                    "stopSequences": []155                }156            }157            158            response_data = self._make_request(url, payload)159            160            # Extract generated text161            if ('candidates' in response_data and 162                len(response_data['candidates']) > 0 and163                'content' in response_data['candidates'][0] and164                'parts' in response_data['candidates'][0]['content']):165                166                generated_text = response_data['candidates'][0]['content']['parts'][0]['text']167                168                # Try to extract suggested questions from the response169                suggested_questions = self._extract_suggested_questions(generated_text)170                171                return {172                    'text': generated_text,173                    'suggested_questions': suggested_questions,174                    'usage': response_data.get('usageMetadata', {})175                }176            else:177                logger.error(f"Unexpected generation response: {response_data}")178                return {179                    'text': 'عذراً، حدث خطأ في توليد الإجابة.',180                    'suggested_questions': ["ما هي عروض السفر؟", "عنّا", "التاشيرات"]181                }182                183        except Exception as e:184            logger.error(f"Error generating response: {e}")185            return {186                'text': 'عذراً، حدث خطأ مؤقت. يرجى المحاولة مرة أخرى.',187                'suggested_questions': ["ما هي عروض السفر؟", "عنّا", "التاشيرات"]188            }189    190    def _extract_suggested_questions(self, text: str) -> List[str]:191        """Extract suggested questions from generated text."""192        # Default suggestions193        default_suggestions = [194            "ما هي عروض السفر؟",195            "عنّا", 196            "التاشيرات",197            "احجز رحلة"198        ]199        200        # Simple heuristic to find questions in the response201        lines = text.split('\n')202        questions = []203        204        for line in lines:205            line = line.strip()206            if line.endswith('؟') and len(line) < 100:  # Arabic question mark207                questions.append(line)208        209        # Return found questions or defaults210        if questions and len(questions) <= 6:211            return questions[:4]  # Max 4 suggestions212        else:213            return default_suggestions