Team Ai
Apppublic

maikheb/nl2sql

sourceHugging Faceupdated 1y agoView on Hugging Face
0likes
gemini_flash_beta_llm.py61 linesDownload Raw Back to src
1# gemini_flash_beta_llm.py
2
3import os
4import requests
5from dotenv import load_dotenv
6from langchain.llms.base import LLM
7from typing import Optional, List
8
9load_dotenv()
10
11# Default settings for the Flash endpoint
12API_KEY = os.getenv("GEMINI_API_KEY")
13BASE_URL = "https://generativelanguage.googleapis.com/v1beta"
14DEFAULT_MODEL_NAME = "gemini-2.0-flash"
15
16class GeminiFlashBetaLLM(LLM):
17    """LangChain LLM wrapper around the v1beta Flash endpoint."""
18    # Pydantic fields for BaseModel compatibility
19    model_name: str = DEFAULT_MODEL_NAME
20    api_key: str = API_KEY
21
22    def _call(self, prompt: str, stop: Optional[List[str]] = None) -> str:
23        # Ensure API key
24        key = self.api_key or API_KEY
25        if not key:
26            raise ValueError("Set your GEMINI_API_KEY in .env")
27
28        url = f"{BASE_URL}/models/{self.model_name}:generateContent"
29        params = {"key": key}
30        body = {"contents": [{"parts": [{"text": prompt}]}]}
31
32        r = requests.post(url, params=params, json=body)
33        r.raise_for_status()
34        data = r.json()
35
36        # Extract the first candidate
37        candidates = data.get("candidates") or []
38        if not candidates:
39            return ""
40        first = candidates[0]
41
42        # Handle content vs output fields
43        content = first.get("content") if "content" in first else first.get("output")
44        text = ""
45
46        if isinstance(content, dict):
47            # Some models return nested parts
48            parts = content.get("parts") or []
49            for part in parts:
50                # support dict or plain text
51                fragment = part.get("text") if isinstance(part, dict) else part
52                if fragment:
53                    text += fragment
54        elif isinstance(content, str):
55            text = content
56
57        return text.strip()
58
59    @property
60    def _llm_type(self) -> str:
61        return f"gemini-flash-v1beta-{self.model_name}"