-
Notifications
You must be signed in to change notification settings - Fork 18
Expand file tree
/
Copy pathllm_providers.py
More file actions
327 lines (259 loc) · 11.9 KB
/
Copy pathllm_providers.py
File metadata and controls
327 lines (259 loc) · 11.9 KB
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
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
"""
LLM Provider abstraction layer for TTP-Threat-Feeds
Supports multiple LLM backends: LM Studio, Ollama, OpenAI, Claude, Gemini
"""
import os
import json
import requests
from abc import ABC, abstractmethod
from typing import Optional, Dict, Any
class LLMProvider(ABC):
"""Base class for LLM providers"""
def __init__(self, model_name: Optional[str] = None, **kwargs):
self.model_name = model_name
self.config = kwargs
@abstractmethod
def generate(self, system_prompt: str, user_prompt: str, temperature: float = 0.2, max_tokens: int = 5000) -> str:
"""
Generate a response from the LLM
Args:
system_prompt: System message/context
user_prompt: User message/query
temperature: Sampling temperature
max_tokens: Maximum tokens to generate
Returns:
Generated text response
"""
pass
@abstractmethod
def get_provider_name(self) -> str:
"""Return the name of this provider"""
pass
class LMStudioProvider(LLMProvider):
"""LM Studio local LLM provider (OpenAI-compatible endpoint)"""
def __init__(self, endpoint: str = "http://127.0.0.1:1234/v1/chat/completions",
model_name: str = "qwen2.5-coder-32b-instruct", **kwargs):
super().__init__(model_name, **kwargs)
self.endpoint = endpoint
self.headers = {"Content-Type": "application/json"}
def generate(self, system_prompt: str, user_prompt: str, temperature: float = 0.2, max_tokens: int = 5000) -> str:
payload = {
"model": self.model_name,
"messages": [
{"role": "system", "content": system_prompt},
{"role": "user", "content": user_prompt}
],
"temperature": temperature,
"max_tokens": max_tokens
}
try:
resp = requests.post(self.endpoint, headers=self.headers, json=payload, timeout=120)
resp.raise_for_status()
data = resp.json()
if "choices" not in data:
raise ValueError(f"Unexpected LLM response: {data}")
return data["choices"][0]["message"]["content"]
except Exception as e:
raise RuntimeError(f"LM Studio API request failed: {e}")
def get_provider_name(self) -> str:
return "LM Studio"
class OllamaProvider(LLMProvider):
"""Ollama local LLM provider"""
def __init__(self, endpoint: str = "http://127.0.0.1:11434/api/chat",
model_name: str = "qwen2.5-coder:32b", **kwargs):
super().__init__(model_name, **kwargs)
self.endpoint = endpoint
self.headers = {"Content-Type": "application/json"}
def generate(self, system_prompt: str, user_prompt: str, temperature: float = 0.2, max_tokens: int = 5000) -> str:
payload = {
"model": self.model_name,
"messages": [
{"role": "system", "content": system_prompt},
{"role": "user", "content": user_prompt}
],
"stream": False,
"options": {
"temperature": temperature,
"num_predict": max_tokens
}
}
try:
resp = requests.post(self.endpoint, headers=self.headers, json=payload, timeout=120)
resp.raise_for_status()
data = resp.json()
if "message" not in data:
raise ValueError(f"Unexpected Ollama response: {data}")
return data["message"]["content"]
except Exception as e:
raise RuntimeError(f"Ollama API request failed: {e}")
def get_provider_name(self) -> str:
return "Ollama"
class OpenAIProvider(LLMProvider):
"""OpenAI API provider"""
def __init__(self, api_key: Optional[str] = None,
model_name: str = "gpt-4o", **kwargs):
super().__init__(model_name, **kwargs)
self.api_key = api_key or os.getenv("OPENAI_API_KEY")
if not self.api_key:
raise ValueError("OpenAI API key not provided. Set OPENAI_API_KEY environment variable or pass api_key parameter")
self.endpoint = "https://api.openai.com/v1/chat/completions"
self.headers = {
"Content-Type": "application/json",
"Authorization": f"Bearer {self.api_key}"
}
def generate(self, system_prompt: str, user_prompt: str, temperature: float = 0.2, max_tokens: int = 5000) -> str:
payload = {
"model": self.model_name,
"messages": [
{"role": "system", "content": system_prompt},
{"role": "user", "content": user_prompt}
],
"temperature": temperature,
"max_tokens": max_tokens
}
try:
resp = requests.post(self.endpoint, headers=self.headers, json=payload, timeout=120)
# If error, try to get the error message
if resp.status_code != 200:
try:
error_data = resp.json()
error_msg = error_data.get('error', {}).get('message', resp.text)
except Exception:
error_msg = resp.text
raise RuntimeError(f"OpenAI API error ({resp.status_code}): {error_msg}")
data = resp.json()
if "choices" not in data:
raise ValueError(f"Unexpected OpenAI response: {data}")
return data["choices"][0]["message"]["content"]
except requests.exceptions.RequestException as e:
raise RuntimeError(f"OpenAI API request failed: {e}")
except Exception as e:
raise RuntimeError(f"OpenAI API error: {e}")
def get_provider_name(self) -> str:
return "OpenAI"
class ClaudeProvider(LLMProvider):
"""Anthropic Claude API provider"""
def __init__(self, api_key: Optional[str] = None,
model_name: str = "claude-sonnet-4-6", **kwargs):
super().__init__(model_name, **kwargs)
self.api_key = api_key or os.getenv("ANTHROPIC_API_KEY")
if not self.api_key:
raise ValueError("Anthropic API key not provided. Set ANTHROPIC_API_KEY environment variable or pass api_key parameter")
self.endpoint = "https://api.anthropic.com/v1/messages"
self.headers = {
"Content-Type": "application/json",
"x-api-key": self.api_key,
"anthropic-version": "2023-06-01"
}
def generate(self, system_prompt: str, user_prompt: str, temperature: float = 0.2, max_tokens: int = 5000) -> str:
# Validate inputs
if not isinstance(system_prompt, str):
raise ValueError(f"system_prompt must be a string, got {type(system_prompt)}")
if not isinstance(user_prompt, str):
raise ValueError(f"user_prompt must be a string, got {type(user_prompt)}")
payload = {
"model": self.model_name,
"system": system_prompt,
"messages": [
{"role": "user", "content": user_prompt}
],
"temperature": temperature,
"max_tokens": max_tokens
}
try:
resp = requests.post(self.endpoint, headers=self.headers, json=payload, timeout=120)
# If error, try to get the error message from Claude's response
if resp.status_code != 200:
try:
error_data = resp.json()
error_msg = error_data.get('error', {}).get('message', resp.text)
except Exception:
error_msg = resp.text
raise RuntimeError(f"Claude API error ({resp.status_code}): {error_msg}")
data = resp.json()
if "content" not in data:
raise ValueError(f"Unexpected Claude response: {data}")
# Claude returns content as a list of content blocks
return data["content"][0]["text"]
except requests.exceptions.RequestException as e:
raise RuntimeError(f"Claude API request failed: {e}")
except Exception as e:
raise RuntimeError(f"Claude API error: {e}")
def get_provider_name(self) -> str:
return "Claude"
class GeminiProvider(LLMProvider):
"""Google Gemini API provider"""
def __init__(self, api_key: Optional[str] = None,
model_name: str = "gemini-2.0-flash-exp", **kwargs):
super().__init__(model_name, **kwargs)
self.api_key = api_key or os.getenv("GOOGLE_API_KEY")
if not self.api_key:
raise ValueError("Google API key not provided. Set GOOGLE_API_KEY environment variable or pass api_key parameter")
self.model_name = model_name
self.endpoint = f"https://generativelanguage.googleapis.com/v1beta/models/{model_name}:generateContent?key={self.api_key}"
self.headers = {"Content-Type": "application/json"}
self.last_request_time = 0 # Track last request for rate limiting
self.min_request_interval = 4.5 # Minimum seconds between requests (15 RPM = 4s, add buffer)
def generate(self, system_prompt: str, user_prompt: str, temperature: float = 0.2, max_tokens: int = 5000) -> str:
import time
# Rate limiting for free tier (15 requests/minute)
current_time = time.time()
time_since_last = current_time - self.last_request_time
if time_since_last < self.min_request_interval:
wait_time = self.min_request_interval - time_since_last
print(f" ⏱️ Gemini rate limit: waiting {wait_time:.1f}s...")
time.sleep(wait_time)
self.last_request_time = time.time()
# Gemini combines system and user prompts differently
combined_prompt = f"{system_prompt}\n\n{user_prompt}"
payload = {
"contents": [{
"parts": [{"text": combined_prompt}]
}],
"generationConfig": {
"temperature": temperature,
"maxOutputTokens": max_tokens,
}
}
try:
resp = requests.post(self.endpoint, headers=self.headers, json=payload, timeout=120)
# Better error handling for Gemini
if resp.status_code == 429:
raise RuntimeError("Gemini rate limit exceeded. Free tier allows 15 requests/minute. Try again in a minute or upgrade your plan.")
elif resp.status_code != 200:
try:
error_data = resp.json()
error_msg = error_data.get('error', {}).get('message', resp.text)
raise ValueError(f"Gemini API error ({resp.status_code}): {error_msg}")
except (ValueError, KeyError):
resp.raise_for_status()
data = resp.json()
if "candidates" not in data or len(data["candidates"]) == 0:
raise ValueError(f"Unexpected Gemini response: {data}")
return data["candidates"][0]["content"]["parts"][0]["text"]
except requests.exceptions.RequestException as e:
raise RuntimeError(f"Gemini API request failed: {e}")
except Exception as e:
raise RuntimeError(f"Gemini API error: {e}")
def get_provider_name(self) -> str:
return "Gemini"
def create_provider(provider_type: str, **kwargs) -> LLMProvider:
"""
Factory function to create LLM providers
Args:
provider_type: One of 'lmstudio', 'ollama', 'openai', 'claude', 'gemini'
**kwargs: Provider-specific configuration
Returns:
LLMProvider instance
"""
providers = {
'lmstudio': LMStudioProvider,
'ollama': OllamaProvider,
'openai': OpenAIProvider,
'claude': ClaudeProvider,
'gemini': GeminiProvider
}
provider_type = provider_type.lower()
if provider_type not in providers:
raise ValueError(f"Unknown provider type: {provider_type}. Available: {', '.join(providers.keys())}")
return providers[provider_type](**kwargs)