File size: 4,888 Bytes
e9b3659 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 | from abc import ABC, abstractmethod
from config import get_settings
class BaseLLM(ABC):
@abstractmethod
def generate(self, prompt: str, system: str = "") -> str:
...
@abstractmethod
def stream(self, prompt: str, system: str = ""):
...
# ββ Gemini provider ββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ
class GeminiLLM(BaseLLM):
def __init__(self, api_key: str, model: str):
from google import genai
self.client = genai.Client(api_key=api_key)
self.model = model
def generate(self, prompt: str, system: str = "") -> str:
full_prompt = f"{system}\n\n{prompt}" if system else prompt
response = self.client.models.generate_content(
model=self.model,
contents=full_prompt,
)
return response.text
def stream(self, prompt: str, system: str = ""):
full_prompt = f"{system}\n\n{prompt}" if system else prompt
for chunk in self.client.models.generate_content_stream(
model=self.model,
contents=full_prompt,
):
if chunk.text:
yield chunk.text
# ββ Groq provider (free tier, no daily quota issues) ββββββββββββββββββββββββββ
class GroqLLM(BaseLLM):
"""
Uses Groq's free API β no per-day quota, just a per-minute rate limit
that is far more generous than Gemini's free tier.
Sign up at https://console.groq.com and set GROQ_API_KEY in your .env.
Recommended model: llama-3.3-70b-versatile (free)
"""
def __init__(self, api_key: str, model: str):
from groq import Groq
self.client = Groq(api_key=api_key)
self.model = model
def _messages(self, prompt: str, system: str) -> list[dict]:
messages = []
if system:
messages.append({"role": "system", "content": system})
messages.append({"role": "user", "content": prompt})
return messages
def generate(self, prompt: str, system: str = "") -> str:
response = self.client.chat.completions.create(
model=self.model,
messages=self._messages(prompt, system),
)
return response.choices[0].message.content
def stream(self, prompt: str, system: str = ""):
stream = self.client.chat.completions.create(
model=self.model,
messages=self._messages(prompt, system),
stream=True,
)
for chunk in stream:
delta = chunk.choices[0].delta.content
if delta:
yield delta
# ββ Factory βββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ
def get_llm_client() -> BaseLLM:
settings = get_settings()
if settings.llm_provider == "gemini":
return GeminiLLM(settings.gemini_api_key, settings.llm_model)
if settings.llm_provider == "groq":
return GroqLLM(settings.groq_api_key, settings.groq_model)
raise ValueError(
f"Unsupported LLM_PROVIDER '{settings.llm_provider}'. "
"Valid options: 'gemini', 'groq'"
)
# ββ Prompt templates ββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ
def build_prompt(query: str, context: str, task_type: str) -> str:
templates = {
"qa": (
"You are a codebase assistant. Use the context below to answer the question.\n\n"
"Context:\n{context}\n\nQuestion: {query}\n\nAnswer clearly, citing file names when relevant."
),
"bug_finding": (
"Analyze the following code for potential bugs, security issues, or bad practices.\n\n"
"Code:\n{context}\n\n"
"Return a JSON list of objects with keys: line, issue, severity, suggestion."
),
"docstring": (
"Write a clear, concise docstring for the following function. "
"Follow standard conventions for its language.\n\nFunction:\n{query}"
),
"file_creation": (
"Generate complete, production-ready code for the following request.\n\n"
"Request: {query}\n\nRelevant existing code for context:\n{context}\n\n"
"Return only the code, no explanations."
),
}
template = templates.get(task_type, templates["qa"])
return template.format(query=query, context=context)
def count_tokens(text: str) -> int:
import tiktoken
encoder = tiktoken.get_encoding("cl100k_base")
return len(encoder.encode(text)) |