llm-proxy-cli 0.5.0__tar.gz → 0.5.2__tar.gz

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.4
2
2
  Name: llm-proxy-cli
3
- Version: 0.5.0
3
+ Version: 0.5.2
4
4
  Summary: A lightweight CLI tool for delegating LLM tasks to expert models across multiple providers.
5
5
  Author-email: Kerem Barbaros Karnabat <kbarbaros@hotmail.com>
6
6
  Classifier: Programming Language :: Python :: 3
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.4
2
2
  Name: llm-proxy-cli
3
- Version: 0.5.0
3
+ Version: 0.5.2
4
4
  Summary: A lightweight CLI tool for delegating LLM tasks to expert models across multiple providers.
5
5
  Author-email: Kerem Barbaros Karnabat <kbarbaros@hotmail.com>
6
6
  Classifier: Programming Language :: Python :: 3
@@ -9,7 +9,7 @@ import logging
9
9
  from openai import OpenAI
10
10
  from filelock import FileLock, Timeout
11
11
 
12
- __version__ = "0.5.0"
12
+ __version__ = "0.5.2"
13
13
 
14
14
  # Optional import for anthropic
15
15
  try:
@@ -48,11 +48,9 @@ def get_api_key(provider):
48
48
  with open(keys_file, 'r', encoding='utf-8') as f:
49
49
  keys = json.load(f)
50
50
  env_name = f"{provider.upper()}_API_KEY"
51
- if keys.get(env_name):
52
- return keys[env_name]
53
51
  except Exception:
54
52
  pass
55
- return os.environ.get(f"{provider.upper()}_API_KEY")
53
+ return os.environ.get(f"{provider.upper()}_API_KEY") or (keys.get(env_name) if 'keys' in locals() else None)
56
54
 
57
55
 
58
56
  PROVIDERS = {
@@ -99,7 +97,11 @@ def _is_chat_candidate(model_id):
99
97
  def _score_model_smart(model_id):
100
98
  """Higher score = bigger / newer / more capable. Used by auto-smart."""
101
99
  name = model_id.lower()
102
- # Strip dates like 2024-08-06 or context sizes like 32768
100
+ # Handle MoE like 8x7b -> 56b
101
+ def _moe(m):
102
+ return str(float(m.group(1)) * float(m.group(2))) + 'b'
103
+ name = re.sub(r'(\d+(?:\.\d+)?)\s*x\s*(\d+(?:\.\d+)?)\s*b', _moe, name)
104
+ # Strip dates and huge context
103
105
  name = re.sub(r'20\d{2}[-]?\d{2}[-]?\d{2}', '', name)
104
106
  name = re.sub(r'\d{4,}', '', name)
105
107
  score = 0.0
@@ -108,35 +110,36 @@ def _score_model_smart(model_id):
108
110
  elif "sonnet" in name: score += 40
109
111
  elif "haiku" in name: score += 30
110
112
 
111
- # Parameter size, e.g. "8b", "70b", "405b" -> biggest wins, weighted heavily.
113
+ # Version extraction (handle 3-5 as 3.5 for claude/llama)
114
+ name_wo_size = re.sub(r'\d+(?:\.\d+)?\s*b(?!\w)', '', name)
115
+ ver_str = name_wo_size.replace('_', '-').replace('-r', '-')
116
+ ver_matches = re.findall(r'(\d+(?:[-.]\d+)*)', ver_str)
117
+ best_ver = 0.0
118
+ for v in ver_matches:
119
+ parts = v.replace('-', '.').split('.')
120
+ val = float(f"{parts[0]}.{parts[1]}") if len(parts) >= 2 else (float(parts[0]) if parts[0] else 0.0)
121
+ best_ver = max(best_ver, val)
122
+ score += best_ver * 5
123
+
124
+ # Parameter size
112
125
  size_matches = re.findall(r'(\d+(?:\.\d+)?)\s*b(?!\w)', name)
113
126
  if size_matches:
114
127
  score += max(float(s) for s in size_matches) * 10
115
- name_wo_size = re.sub(r'\d+(?:\.\d+)?\s*b(?!\w)', '', name)
116
- else:
117
- name_wo_size = name
118
-
119
- # Any remaining digits are treated as generation/version numbers,
120
- # e.g. "llama-4", "gemini-2.5", "v3.3" -> higher wins.
121
- version_matches = re.findall(r'(\d+(?:\.\d+)?)', name_wo_size)
122
- if version_matches:
123
- score += max(float(v) for v in version_matches) * 5
124
128
 
125
129
  for kw, bonus in (("instruct", 3), ("versatile", 3), ("reasoning", 4), ("chat", 1)):
126
- if kw in name:
127
- score += bonus
130
+ if kw in name: score += bonus
128
131
  for kw, penalty in (("mini", -2), ("lite", -2), ("tiny", -3), ("preview", -1), ("deprecated", -100)):
129
- if kw in name:
130
- score += penalty
132
+ if kw in name: score += penalty
131
133
 
132
134
  return score
133
135
 
134
136
 
135
137
  def _score_model_fast(model_id):
136
- """Higher score = smaller / snappier chat model. Used by auto-fast.
137
- NON_CHAT_KEYWORDS models are filtered out before this ever runs, so
138
- "smallest wins" can't land on a non-chat model anymore."""
138
+ """Higher score = smaller / snappier chat model. Used by auto-fast."""
139
139
  name = model_id.lower()
140
+ def _moe(m):
141
+ return str(float(m.group(1)) * float(m.group(2))) + 'b'
142
+ name = re.sub(r'(\d+(?:\.\d+)?)\s*x\s*(\d+(?:\.\d+)?)\s*b', _moe, name)
140
143
  name = re.sub(r'20\d{2}[-]?\d{2}[-]?\d{2}', '', name)
141
144
  name = re.sub(r'\d{4,}', '', name)
142
145
  score = 0.0
@@ -144,23 +147,24 @@ def _score_model_fast(model_id):
144
147
  if "haiku" in name: score += 20
145
148
  elif "sonnet" in name: score += 10
146
149
 
150
+ name_wo_size = re.sub(r'\d+(?:\.\d+)?\s*b(?!\w)', '', name)
151
+ ver_str = name_wo_size.replace('_', '-').replace('-r', '-')
152
+ ver_matches = re.findall(r'(\d+(?:[-.]\d+)*)', ver_str)
153
+ best_ver = 0.0
154
+ for v in ver_matches:
155
+ parts = v.replace('-', '.').split('.')
156
+ val = float(f"{parts[0]}.{parts[1]}") if len(parts) >= 2 else (float(parts[0]) if parts[0] else 0.0)
157
+ best_ver = max(best_ver, val)
158
+ score += best_ver * 5
159
+
147
160
  size_matches = re.findall(r'(\d+(?:\.\d+)?)\s*b(?!\w)', name)
148
161
  if size_matches:
149
- score -= max(float(s) for s in size_matches) * 10 # smaller size = higher score
150
- name_wo_size = re.sub(r'\d+(?:\.\d+)?\s*b(?!\w)', '', name)
151
- else:
152
- name_wo_size = name
153
-
154
- version_matches = re.findall(r'(\d+(?:\.\d+)?)', name_wo_size)
155
- if version_matches:
156
- score += max(float(v) for v in version_matches) * 5 # still prefer the newer generation
162
+ score -= max(float(s) for s in size_matches) * 10
157
163
 
158
164
  for kw, bonus in (("instant", 4), ("flash", 4), ("turbo", 4), ("mini", 3), ("lite", 3), ("small", 2), ("instruct", 1)):
159
- if kw in name:
160
- score += bonus
165
+ if kw in name: score += bonus
161
166
  for kw, penalty in (("preview", -1), ("deprecated", -100)):
162
- if kw in name:
163
- score += penalty
167
+ if kw in name: score += penalty
164
168
 
165
169
  return score
166
170
 
@@ -332,8 +336,16 @@ class CircuitBreaker:
332
336
  with FileLock(self.lock_file, timeout=5):
333
337
  circuit = self.load()
334
338
  if model_id not in circuit:
335
- circuit[model_id] = {'failures': 0, 'cooldown_until': 0}
339
+ circuit[model_id] = {'failures': 0, 'cooldown_until': 0, 'last_failure_time': 0}
340
+
341
+ # Decay old failures after 300 seconds
342
+ last_fail = circuit[model_id].get('last_failure_time', 0)
343
+ if time.time() - last_fail > 300:
344
+ circuit[model_id]['failures'] = 0
345
+
336
346
  circuit[model_id]['failures'] += 1
347
+ circuit[model_id]['last_failure_time'] = time.time()
348
+
337
349
  if circuit[model_id]['failures'] >= self.max_failures:
338
350
  circuit[model_id]['cooldown_until'] = time.time() + self.cooldown_seconds
339
351
  circuit[model_id]['failures'] = 0
@@ -370,17 +382,19 @@ def parse_model(model_string, force_refresh_auto=False):
370
382
  if resolved:
371
383
  return provider, resolved
372
384
  logger.error(f"Auto-discovery unavailable for '{provider}' ({mode}) and no static fallback was given.")
373
- return provider, model # will fail fast in query_ai (model not found / no key)
385
+ return provider, None
374
386
 
375
387
  return provider, model
376
388
 
377
389
 
378
- def query_ai(models_list, prompt, cb: CircuitBreaker, max_retries=2, base_timeout=30, force_refresh_auto=False):
390
+ def query_ai(models_list, prompt, cb: CircuitBreaker, max_retries=2, base_timeout=30, force_refresh_auto=False, stream_out=False):
379
391
  if isinstance(models_list, str):
380
392
  models_list = [m.strip() for m in models_list.split(',')]
381
393
 
382
394
  for current_model_str in models_list:
383
395
  provider, current_model = parse_model(current_model_str, force_refresh_auto=force_refresh_auto)
396
+ if not current_model:
397
+ continue
384
398
  # Circuit breaker tracks the *resolved* model, not the literal "auto" alias,
385
399
  # since "auto" can point at a different real model over time.
386
400
  resolved_key = f"{provider}:{current_model}"
@@ -410,6 +424,9 @@ def query_ai(models_list, prompt, cb: CircuitBreaker, max_retries=2, base_timeou
410
424
  messages=[{"role": "user", "content": prompt}], temperature=0.7
411
425
  ) as stream:
412
426
  for text in stream.text_stream:
427
+ if stream_out:
428
+ sys.stdout.write(text)
429
+ sys.stdout.flush()
413
430
  full_content += text
414
431
  else:
415
432
  client = OpenAI(
@@ -423,9 +440,17 @@ def query_ai(models_list, prompt, cb: CircuitBreaker, max_retries=2, base_timeou
423
440
  for chunk in completion:
424
441
  if not chunk.choices: continue
425
442
  reasoning = getattr(chunk.choices[0].delta, "reasoning_content", None)
426
- if reasoning: full_reasoning += reasoning
443
+ if reasoning:
444
+ full_reasoning += reasoning
445
+ if stream_out:
446
+ sys.stderr.write(reasoning)
447
+ sys.stderr.flush()
427
448
  content = chunk.choices[0].delta.content
428
- if content: full_content += content
449
+ if content:
450
+ if stream_out:
451
+ sys.stdout.write(content)
452
+ sys.stdout.flush()
453
+ full_content += content
429
454
 
430
455
  cb.record_success(resolved_key)
431
456
  output = ""
@@ -436,10 +461,15 @@ def query_ai(models_list, prompt, cb: CircuitBreaker, max_retries=2, base_timeou
436
461
  except Exception as e:
437
462
  error_msg = str(e).lower()
438
463
  logger.error(f"Attempt {attempt+1} failed for {resolved_key}: {str(e)}")
439
- cb.record_failure(resolved_key)
440
-
464
+ if stream_out and (full_content or full_reasoning):
465
+ sys.stderr.write(f"\n[STREAM INTERRUPTED: {str(e)} - FALLBACK TRIGGERED]\n")
466
+ sys.stderr.flush()
441
467
  status_code = getattr(e, "status_code", None)
442
- if status_code in (429, 404, 401, 403) or "404" in error_msg or "not found" in error_msg or "auth" in error_msg:
468
+ if status_code in (404, 401, 403) or "404" in error_msg or "not found" in error_msg or "auth" in error_msg:
469
+ break
470
+ if status_code != 429:
471
+ cb.record_failure(resolved_key)
472
+ if status_code == 429:
443
473
  break
444
474
 
445
475
  if attempt == max_retries - 1: break
@@ -452,8 +482,8 @@ def main():
452
482
  if hasattr(sys.stdout, 'reconfigure'):
453
483
  sys.stdout.reconfigure(encoding='utf-8')
454
484
 
455
- parser = argparse.ArgumentParser(description="Smart Router: A fault-tolerant CLI tool for LLM delegation.")
456
- parser.add_argument("-v", "--version", action="version", version=f"Smart Router v{__version__}")
485
+ parser = argparse.ArgumentParser(description="LLM Proxy CLI: A fault-tolerant CLI tool for LLM delegation.")
486
+ parser.add_argument("-v", "--version", action="version", version=f"LLM Proxy CLI v{__version__}")
457
487
  parser.add_argument("-m", "--models", required=True, help="Comma-separated list of provider:model fallbacks (e.g. nvidia:nemotron,groq:llama3, or groq:auto-smart / groq:auto-fast).")
458
488
  parser.add_argument("-p", "--prompt", help="The prompt text to send to the model.")
459
489
  parser.add_argument("-f", "--file", help="Path to a text file containing the prompt.")
@@ -484,8 +514,8 @@ def main():
484
514
 
485
515
  cb = CircuitBreaker(args.project, args.max_failures, args.cooldown)
486
516
  try:
487
- response = query_ai(args.models, prompt_text, cb, force_refresh_auto=args.refresh_models)
488
- print(response)
517
+ query_ai(args.models, prompt_text, cb, force_refresh_auto=args.refresh_models, stream_out=True)
518
+ print()
489
519
  except Exception as e:
490
520
  logger.error(str(e))
491
521
  sys.exit(1)
@@ -4,7 +4,7 @@ build-backend = "setuptools.build_meta"
4
4
 
5
5
  [project]
6
6
  name = "llm-proxy-cli"
7
- version = "0.5.0"
7
+ version = "0.5.2"
8
8
  authors = [
9
9
  { name="Kerem Barbaros Karnabat", email="kbarbaros@hotmail.com" }
10
10
  ]
@@ -46,3 +46,46 @@ def test_circuit_breaker_success_reset(tmp_path):
46
46
  cb.record_success("modelC")
47
47
  circuit = cb.load()
48
48
  assert "modelC" not in circuit or circuit["modelC"]["failures"] == 0
49
+
50
+ import time
51
+
52
+ def test_circuit_breaker_decay():
53
+ cb = CircuitBreaker("decay_test", max_failures=2, cooldown_seconds=60)
54
+
55
+ # Clean state
56
+ if os.path.exists(cb.circuit_file):
57
+ os.remove(cb.circuit_file)
58
+
59
+ import llm_proxy_cli
60
+
61
+ # Mock time.time to simulate passage of time
62
+ original_time = time.time
63
+ current_mock_time = original_time()
64
+
65
+ def mock_time():
66
+ return current_mock_time
67
+
68
+ llm_proxy_cli.time.time = mock_time
69
+
70
+ try:
71
+ # Failure 1 at T=0
72
+ cb.record_failure("test:model")
73
+ assert cb.check_health("test:model") == True
74
+
75
+ # Advance time by 400 seconds (past 300s decay threshold)
76
+ current_mock_time += 400
77
+
78
+ # Failure 2 at T=400
79
+ # This should reset the counter to 0 before adding 1, so it won't trip (max_failures=2)
80
+ cb.record_failure("test:model")
81
+ assert cb.check_health("test:model") == True
82
+
83
+ # Failure 3 immediately after (T=400)
84
+ # Counter becomes 2 -> trips!
85
+ cb.record_failure("test:model")
86
+ assert cb.check_health("test:model") == False
87
+
88
+ finally:
89
+ llm_proxy_cli.time.time = original_time
90
+ if os.path.exists(cb.circuit_file):
91
+ os.remove(cb.circuit_file)
@@ -117,7 +117,7 @@ def test_query_ai_stops_retrying_immediately_on_404(monkeypatch, mock_cb):
117
117
  assert result == "ok"
118
118
  # A 404 must break the retry loop after a single attempt, not consume all max_retries=3.
119
119
  assert not_found_client.chat.completions.create.call_count == 1
120
- assert mock_cb.record_failure.call_count == 1
120
+ assert mock_cb.record_failure.call_count == 0
121
121
 
122
122
 
123
123
  def test_query_ai_skips_models_in_cooldown(monkeypatch, mock_cb):
@@ -159,10 +159,10 @@ def test_query_ai_exits_when_every_model_fails(monkeypatch, mock_cb):
159
159
  failing_client.chat.completions.create.side_effect = Exception("timeout")
160
160
  factory = openai_factory({GROQ_URL: failing_client})
161
161
  with patch.object(router, "OpenAI", side_effect=factory):
162
- with pytest.raises(SystemExit) as exc_info:
162
+ with pytest.raises(RuntimeError) as exc_info:
163
163
  router.query_ai("groq:llama-3.3-70b-versatile", "hi", mock_cb, max_retries=2)
164
164
 
165
- assert exc_info.value.code == 1
165
+ assert "All fallback models failed" in str(exc_info.value)
166
166
  assert failing_client.chat.completions.create.call_count == 2
167
167
 
168
168
 
@@ -197,3 +197,38 @@ def test_query_ai_resolves_auto_alias_before_calling_provider(monkeypatch, mock_
197
197
  # circuit breaker must key on the resolved model, never on the literal "auto-smart" alias
198
198
  mock_cb.check_health.assert_called_with("groq:resolved-model-xyz")
199
199
  mock_cb.record_success.assert_called_once_with("groq:resolved-model-xyz")
200
+
201
+
202
+ def test_query_ai_streaming_output(capsys, monkeypatch, mock_cb):
203
+ """Test that stream_out=True sends reasoning to stderr and content to stdout, exactly once."""
204
+ monkeypatch.setitem(router.PROVIDERS, "groq", {"base_url": "dummy", "api_key": "key"})
205
+
206
+ mock_client = MagicMock()
207
+ mock_chunk1 = MagicMock()
208
+ mock_chunk1.choices = [MagicMock()]
209
+ mock_chunk1.choices[0].delta.reasoning_content = "thinking..."
210
+ mock_chunk1.choices[0].delta.content = ""
211
+
212
+ mock_chunk2 = MagicMock()
213
+ mock_chunk2.choices = [MagicMock()]
214
+ mock_chunk2.choices[0].delta.reasoning_content = ""
215
+ mock_chunk2.choices[0].delta.content = "hello world"
216
+
217
+ mock_client.chat.completions.create.return_value = [mock_chunk1, mock_chunk2]
218
+
219
+ factory = openai_factory({"dummy": mock_client})
220
+ with patch.object(router, "OpenAI", side_effect=factory):
221
+ response = router.query_ai("groq:llama", "hi", mock_cb, stream_out=True)
222
+
223
+ captured = capsys.readouterr()
224
+
225
+ # 1. Stdout must only contain "hello world" (no reasoning)
226
+ assert "hello world" in captured.out
227
+ assert "thinking..." not in captured.out
228
+
229
+ # 2. Stderr must contain the reasoning
230
+ assert "thinking..." in captured.err
231
+
232
+ # 3. The returned string still contains everything for library usage
233
+ assert "hello world" in response
234
+ assert "thinking..." in response
File without changes
File without changes
File without changes