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 # compound routing models share TPM pools — enforce a longer minimum wait 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) # groq/compound has a strict payload limit; truncate to ~8 000 chars to stay under 413 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: # Payload too large for this model — return a sentinel so the # experiment records the case as failed rather than crashing. 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