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))