Commit ·
2fe3d02
1
Parent(s): 15f6b5f
Support NVIDIA and Mistral provider setup
Browse files- mp1/pluto/dispatcher.py +10 -3
- mp1/pluto/modes.py +52 -3
mp1/pluto/dispatcher.py
CHANGED
|
@@ -28,6 +28,10 @@ NVIDIA_RERANK_URLS = (
|
|
| 28 |
"https://ai.api.nvidia.com/v1/retrieval/nvidia/reranking",
|
| 29 |
"https://integrate.api.nvidia.com/v1/ranking",
|
| 30 |
)
|
|
|
|
|
|
|
|
|
|
|
|
|
| 31 |
|
| 32 |
# Groq has strict RPM limits; semaphore prevents parallel-call overflow.
|
| 33 |
_groq_semaphore = threading.BoundedSemaphore(2)
|
|
@@ -234,6 +238,9 @@ def _call_nvidia(cfg: ModeConfig, prompt: str) -> str:
|
|
| 234 |
env_var, api_key = _resolve_nvidia_api_key(cfg.model_id)
|
| 235 |
if not api_key:
|
| 236 |
raise ValueError(f"{env_var} or NVIDIA_API_KEY not set")
|
|
|
|
|
|
|
|
|
|
| 237 |
|
| 238 |
prefix = str(prompt)[:120]
|
| 239 |
use_reasoning = any(
|
|
@@ -250,7 +257,7 @@ def _call_nvidia(cfg: ModeConfig, prompt: str) -> str:
|
|
| 250 |
if "nano" in cfg.model_id and use_reasoning:
|
| 251 |
payload["nvidia"] = {"reasoning": True}
|
| 252 |
|
| 253 |
-
for attempt in range(
|
| 254 |
try:
|
| 255 |
response = requests.post(
|
| 256 |
NVIDIA_CHAT_URL,
|
|
@@ -259,7 +266,7 @@ def _call_nvidia(cfg: ModeConfig, prompt: str) -> str:
|
|
| 259 |
"Content-Type": "application/json",
|
| 260 |
},
|
| 261 |
json=payload,
|
| 262 |
-
timeout=
|
| 263 |
)
|
| 264 |
if response.status_code == 200:
|
| 265 |
data = response.json()
|
|
@@ -289,7 +296,7 @@ def _call_nvidia(cfg: ModeConfig, prompt: str) -> str:
|
|
| 289 |
"504",
|
| 290 |
]
|
| 291 |
)
|
| 292 |
-
if (is_rate or is_transient) and attempt <
|
| 293 |
delay = min(60.0, 5.0 * (1.5**attempt))
|
| 294 |
reason = "rate-limit" if is_rate else "transient error"
|
| 295 |
print(f" [RETRY] NVIDIA {reason} — waiting {delay:.1f}s (attempt {attempt + 1})")
|
|
|
|
| 28 |
"https://ai.api.nvidia.com/v1/retrieval/nvidia/reranking",
|
| 29 |
"https://integrate.api.nvidia.com/v1/ranking",
|
| 30 |
)
|
| 31 |
+
NVIDIA_MAX_RETRIES_WITH_FALLBACK = 2
|
| 32 |
+
NVIDIA_MAX_RETRIES_DEFAULT = 6
|
| 33 |
+
NVIDIA_TIMEOUT_WITH_FALLBACK = (10, 45)
|
| 34 |
+
NVIDIA_TIMEOUT_DEFAULT = (20, 120)
|
| 35 |
|
| 36 |
# Groq has strict RPM limits; semaphore prevents parallel-call overflow.
|
| 37 |
_groq_semaphore = threading.BoundedSemaphore(2)
|
|
|
|
| 238 |
env_var, api_key = _resolve_nvidia_api_key(cfg.model_id)
|
| 239 |
if not api_key:
|
| 240 |
raise ValueError(f"{env_var} or NVIDIA_API_KEY not set")
|
| 241 |
+
has_fallback = bool(os.getenv("GROQ_API_KEY", "") or os.getenv("MISTRAL_API_KEY", ""))
|
| 242 |
+
max_retries = NVIDIA_MAX_RETRIES_WITH_FALLBACK if has_fallback else NVIDIA_MAX_RETRIES_DEFAULT
|
| 243 |
+
request_timeout = NVIDIA_TIMEOUT_WITH_FALLBACK if has_fallback else NVIDIA_TIMEOUT_DEFAULT
|
| 244 |
|
| 245 |
prefix = str(prompt)[:120]
|
| 246 |
use_reasoning = any(
|
|
|
|
| 257 |
if "nano" in cfg.model_id and use_reasoning:
|
| 258 |
payload["nvidia"] = {"reasoning": True}
|
| 259 |
|
| 260 |
+
for attempt in range(max_retries):
|
| 261 |
try:
|
| 262 |
response = requests.post(
|
| 263 |
NVIDIA_CHAT_URL,
|
|
|
|
| 266 |
"Content-Type": "application/json",
|
| 267 |
},
|
| 268 |
json=payload,
|
| 269 |
+
timeout=request_timeout,
|
| 270 |
)
|
| 271 |
if response.status_code == 200:
|
| 272 |
data = response.json()
|
|
|
|
| 296 |
"504",
|
| 297 |
]
|
| 298 |
)
|
| 299 |
+
if (is_rate or is_transient) and attempt < max_retries - 1:
|
| 300 |
delay = min(60.0, 5.0 * (1.5**attempt))
|
| 301 |
reason = "rate-limit" if is_rate else "transient error"
|
| 302 |
print(f" [RETRY] NVIDIA {reason} — waiting {delay:.1f}s (attempt {attempt + 1})")
|
mp1/pluto/modes.py
CHANGED
|
@@ -32,7 +32,7 @@ class ModeConfig:
|
|
| 32 |
temperature: float
|
| 33 |
max_tokens: int
|
| 34 |
compute_profile: str
|
| 35 |
-
provider: str # "groq" | "mistral"
|
| 36 |
|
| 37 |
def to_log_dict(self) -> dict:
|
| 38 |
return {
|
|
@@ -58,7 +58,7 @@ def _build_registry() -> dict[str, ModeConfig]:
|
|
| 58 |
Embedding + reranking are handled separately in embedder.py and dispatcher.py
|
| 59 |
(they use /v1/embeddings and scoring endpoints, not chat completions).
|
| 60 |
|
| 61 |
-
Fallback: if NVIDIA_API_KEY absent, fall back to Groq
|
| 62 |
"""
|
| 63 |
# Check for any NVIDIA key
|
| 64 |
nvidia_keys = [
|
|
@@ -68,6 +68,7 @@ def _build_registry() -> dict[str, ModeConfig]:
|
|
| 68 |
]
|
| 69 |
nvidia_ready = any(os.getenv(k) for k in nvidia_keys)
|
| 70 |
groq_key = os.getenv("GROQ_API_KEY", "").strip()
|
|
|
|
| 71 |
|
| 72 |
if nvidia_ready:
|
| 73 |
return {
|
|
@@ -157,9 +158,57 @@ def _build_registry() -> dict[str, ModeConfig]:
|
|
| 157 |
provider="groq",
|
| 158 |
),
|
| 159 |
}
|
|
|
|
|
|
|
| 160 |
return _build_unconfigured_registry()
|
| 161 |
|
| 162 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 163 |
def _build_unconfigured_registry() -> dict[str, ModeConfig]:
|
| 164 |
"""Return placeholder modes so imports work without provider credentials."""
|
| 165 |
return {
|
|
@@ -210,7 +259,7 @@ MODE_REGISTRY: dict[str, ModeConfig] = _build_registry()
|
|
| 210 |
|
| 211 |
|
| 212 |
def _missing_provider_error() -> EnvironmentError:
|
| 213 |
-
return EnvironmentError("
|
| 214 |
|
| 215 |
|
| 216 |
def _is_unconfigured() -> bool:
|
|
|
|
| 32 |
temperature: float
|
| 33 |
max_tokens: int
|
| 34 |
compute_profile: str
|
| 35 |
+
provider: str # "nvidia" | "groq" | "mistral"
|
| 36 |
|
| 37 |
def to_log_dict(self) -> dict:
|
| 38 |
return {
|
|
|
|
| 58 |
Embedding + reranking are handled separately in embedder.py and dispatcher.py
|
| 59 |
(they use /v1/embeddings and scoring endpoints, not chat completions).
|
| 60 |
|
| 61 |
+
Fallback: if NVIDIA_API_KEY absent, fall back to Groq or Mistral.
|
| 62 |
"""
|
| 63 |
# Check for any NVIDIA key
|
| 64 |
nvidia_keys = [
|
|
|
|
| 68 |
]
|
| 69 |
nvidia_ready = any(os.getenv(k) for k in nvidia_keys)
|
| 70 |
groq_key = os.getenv("GROQ_API_KEY", "").strip()
|
| 71 |
+
mistral_key = os.getenv("MISTRAL_API_KEY", "").strip()
|
| 72 |
|
| 73 |
if nvidia_ready:
|
| 74 |
return {
|
|
|
|
| 158 |
provider="groq",
|
| 159 |
),
|
| 160 |
}
|
| 161 |
+
if mistral_key:
|
| 162 |
+
return _build_mistral_registry()
|
| 163 |
return _build_unconfigured_registry()
|
| 164 |
|
| 165 |
|
| 166 |
+
def _build_mistral_registry() -> dict[str, ModeConfig]:
|
| 167 |
+
"""Use Mistral for every mode when it is the only configured chat provider."""
|
| 168 |
+
return {
|
| 169 |
+
"MODE_QUICK": ModeConfig(
|
| 170 |
+
mode_name="MODE_QUICK",
|
| 171 |
+
model_id="mistral-small-latest",
|
| 172 |
+
temperature=0.1,
|
| 173 |
+
max_tokens=1024,
|
| 174 |
+
compute_profile="fallback",
|
| 175 |
+
provider="mistral",
|
| 176 |
+
),
|
| 177 |
+
"MODE_REASONING": ModeConfig(
|
| 178 |
+
mode_name="MODE_REASONING",
|
| 179 |
+
model_id="mistral-small-latest",
|
| 180 |
+
temperature=0.3,
|
| 181 |
+
max_tokens=4096,
|
| 182 |
+
compute_profile="fallback",
|
| 183 |
+
provider="mistral",
|
| 184 |
+
),
|
| 185 |
+
"MODE_VISION": ModeConfig(
|
| 186 |
+
mode_name="MODE_VISION",
|
| 187 |
+
model_id="mistral-small-latest",
|
| 188 |
+
temperature=0.1,
|
| 189 |
+
max_tokens=4096,
|
| 190 |
+
compute_profile="fallback",
|
| 191 |
+
provider="mistral",
|
| 192 |
+
),
|
| 193 |
+
"MODE_ULTRA": ModeConfig(
|
| 194 |
+
mode_name="MODE_ULTRA",
|
| 195 |
+
model_id="mistral-small-latest",
|
| 196 |
+
temperature=0.2,
|
| 197 |
+
max_tokens=4096,
|
| 198 |
+
compute_profile="fallback",
|
| 199 |
+
provider="mistral",
|
| 200 |
+
),
|
| 201 |
+
"MODE_GEMINI": ModeConfig(
|
| 202 |
+
mode_name="MODE_GEMINI",
|
| 203 |
+
model_id="mistral-small-latest",
|
| 204 |
+
temperature=0.0,
|
| 205 |
+
max_tokens=4096,
|
| 206 |
+
compute_profile="fallback",
|
| 207 |
+
provider="mistral",
|
| 208 |
+
),
|
| 209 |
+
}
|
| 210 |
+
|
| 211 |
+
|
| 212 |
def _build_unconfigured_registry() -> dict[str, ModeConfig]:
|
| 213 |
"""Return placeholder modes so imports work without provider credentials."""
|
| 214 |
return {
|
|
|
|
| 259 |
|
| 260 |
|
| 261 |
def _missing_provider_error() -> EnvironmentError:
|
| 262 |
+
return EnvironmentError("None of NVIDIA_API_KEY, GROQ_API_KEY, or MISTRAL_API_KEY is set.")
|
| 263 |
|
| 264 |
|
| 265 |
def _is_unconfigured() -> bool:
|