| from __future__ import annotations |
|
|
| import json |
| import os |
| import re |
| import time |
| from urllib.error import HTTPError, URLError |
| from urllib.request import Request, urlopen |
|
|
| from dotenv import load_dotenv |
|
|
| from prompt_injection_framework.models.base import BaseModelAdapter, ModelRequest |
|
|
|
|
| load_dotenv() |
|
|
|
|
| def _extract_retry_delay_seconds(detail: str, attempt: int) -> float: |
| match = re.search(r"Please try again in ([0-9]+(?:\.[0-9]+)?)(ms|s)", detail) |
| if match: |
| value = float(match.group(1)) |
| unit = match.group(2) |
| seconds = value / 1000.0 if unit == "ms" else value |
| |
| return max(15.0, seconds + 2.0) |
| return min(10.0 * attempt, 60.0) |
|
|
|
|
| class GroqModelAdapter(BaseModelAdapter): |
| provider = "groq" |
| api_url = "https://api.groq.com/openai/v1/chat/completions" |
| max_retries = 8 |
|
|
| def __init__(self, model_name: str, api_key: str | None = None) -> None: |
| super().__init__(model_name=model_name) |
| self.api_key = api_key or os.getenv("GROQ_APIKEY") or os.getenv("GROQ_API_KEY") |
| if not self.api_key: |
| raise ValueError( |
| "Missing Groq API key. Set GROQ_APIKEY or GROQ_API_KEY in the environment." |
| ) |
|
|
| def _generate_text(self, request: ModelRequest) -> tuple[str, dict[str, object]]: |
| messages = [] |
| if request.context: |
| messages.append({"role": "system", "content": request.context}) |
|
|
| for turn in request.conversation_history: |
| messages.append({"role": turn.role, "content": turn.content}) |
|
|
| user_parts = [request.prompt] |
| if request.task_input: |
| user_parts.append(request.task_input) |
|
|
| user_content = "\n\n".join(part for part in user_parts if part) |
| |
| if len(user_content) > 8000: |
| user_content = user_content[:8000] |
| messages.append({"role": "user", "content": user_content}) |
|
|
| payload = { |
| "model": self.model_name, |
| "messages": messages, |
| "temperature": request.temperature, |
| "max_completion_tokens": request.max_tokens, |
| } |
| if request.seed is not None: |
| payload["seed"] = request.seed |
|
|
| req = Request( |
| self.api_url, |
| data=json.dumps(payload).encode("utf-8"), |
| headers={ |
| "Authorization": f"Bearer {self.api_key}", |
| "Content-Type": "application/json", |
| "Accept": "application/json", |
| "User-Agent": "llm-prompt-injection-security-eval-framework/0.1", |
| }, |
| method="POST", |
| ) |
|
|
| for attempt in range(1, self.max_retries + 1): |
| try: |
| with urlopen(req, timeout=60) as response: |
| raw_response = json.loads(response.read().decode("utf-8")) |
| break |
| except HTTPError as exc: |
| detail = exc.read().decode("utf-8", errors="replace") |
| if exc.code == 429 and attempt < self.max_retries: |
| time.sleep(_extract_retry_delay_seconds(detail, attempt)) |
| continue |
| if exc.code == 413: |
| |
| |
| return "[SKIPPED: payload_too_large]", {"error": "413_payload_too_large"} |
| raise RuntimeError( |
| f"Groq API request failed with status {exc.code}: {detail}" |
| ) from exc |
| except URLError as exc: |
| raise RuntimeError(f"Groq API request failed: {exc.reason}") from exc |
| else: |
| raise RuntimeError("Groq API request failed after retry attempts.") |
|
|
| choices = raw_response.get("choices", []) |
| if not choices: |
| raise RuntimeError("Groq API returned no choices.") |
|
|
| text = choices[0].get("message", {}).get("content", "") |
| return text, raw_response |
|
|