maikheb/nl2sql
0
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}"