mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-15 23:31:29 +00:00
Encode complete non-Claude source names and include source_model in the Claude Code listing. Preserve configured route and alias precedence, normalize once before model policy checks, and select CLI models using explicit source identity instead of name stripping or positional joins. Resolves LIT-7360 Claude-Session: https://claude.ai/code/session_01WyqeRhfZGm26zAnHx9P3kq Co-authored-by: Claude Fable 5.1 <noreply@anthropic.com>
294 lines
11 KiB
Python
294 lines
11 KiB
Python
import json
|
|
import os
|
|
import stat
|
|
from concurrent.futures import ThreadPoolExecutor
|
|
from pathlib import Path
|
|
|
|
import pytest
|
|
import requests
|
|
|
|
from litellm.proxy.client.cli.commands.pi import (
|
|
ListingFailure,
|
|
ModelLimits,
|
|
PiSyncError,
|
|
fetch_model_ids,
|
|
fetch_model_limits,
|
|
fetch_model_listing,
|
|
models_json_path,
|
|
provider_block,
|
|
sync_models_json,
|
|
)
|
|
|
|
|
|
class _FakeResponse:
|
|
def __init__(self, status_code, payload=None):
|
|
self.status_code = status_code
|
|
self._payload = payload
|
|
|
|
def json(self):
|
|
if self._payload is None:
|
|
raise ValueError("not json")
|
|
return self._payload
|
|
|
|
|
|
def _refused(*args, **kwargs):
|
|
raise requests.ConnectionError("refused")
|
|
|
|
|
|
class TestFetchModelIds:
|
|
def test_returns_ids_in_proxy_order_deduped(self):
|
|
captured = {}
|
|
|
|
def fake_get(url, headers, timeout):
|
|
captured["url"] = url
|
|
captured["headers"] = headers
|
|
return _FakeResponse(
|
|
200,
|
|
{"data": [{"id": "m-b"}, {"id": "m-a"}, {"id": "m-b"}]},
|
|
)
|
|
|
|
assert fetch_model_ids("http://localhost:4000/", "sk-key", get=fake_get) == ("m-b", "m-a")
|
|
assert captured["url"] == "http://localhost:4000/v1/models"
|
|
assert captured["headers"] == {"Authorization": "Bearer sk-key"}
|
|
|
|
def test_returns_rows_with_optional_source_model_and_dedups_identical_rows(self):
|
|
result = fetch_model_listing(
|
|
"http://localhost:4000",
|
|
"sk-key",
|
|
get=lambda *a, **k: _FakeResponse(
|
|
200,
|
|
{"data": [{"id": "emitted", "source_model": "source"}, {"id": "emitted", "source_model": "source"}]},
|
|
),
|
|
)
|
|
assert not isinstance(result, PiSyncError)
|
|
assert tuple((model.id, model.source_model) for model in result) == (("emitted", "source"),)
|
|
|
|
@pytest.mark.parametrize(
|
|
"entry",
|
|
[
|
|
{"id": ""},
|
|
{"id": "emitted", "source_model": ""},
|
|
{"id": "emitted", "source_model": 1},
|
|
],
|
|
)
|
|
def test_rejects_invalid_model_identity(self, entry):
|
|
result = fetch_model_listing(
|
|
"http://localhost:4000", "sk-key", get=lambda *a, **k: _FakeResponse(200, {"data": [entry]})
|
|
)
|
|
assert isinstance(result, PiSyncError) and result.kind is ListingFailure.BAD_BODY
|
|
|
|
def test_rejects_conflicting_emitted_id_mappings(self):
|
|
result = fetch_model_listing(
|
|
"http://localhost:4000",
|
|
"sk-key",
|
|
get=lambda *a, **k: _FakeResponse(
|
|
200,
|
|
{"data": [{"id": "emitted", "source_model": "one"}, {"id": "emitted", "source_model": "two"}]},
|
|
),
|
|
)
|
|
assert isinstance(result, PiSyncError) and result.kind is ListingFailure.BAD_BODY
|
|
|
|
def test_network_error_is_a_value(self):
|
|
def boom(*a, **k):
|
|
raise requests.ConnectionError("refused")
|
|
|
|
result = fetch_model_ids("http://localhost:4000", "sk-key", get=boom)
|
|
assert isinstance(result, PiSyncError)
|
|
assert "Could not list models" in result.message
|
|
|
|
def test_non_200_is_a_value(self):
|
|
result = fetch_model_ids("http://localhost:4000", "sk-key", get=lambda *a, **k: _FakeResponse(500))
|
|
assert isinstance(result, PiSyncError)
|
|
assert "HTTP 500" in result.message
|
|
|
|
def test_malformed_body_is_a_value(self):
|
|
result = fetch_model_ids(
|
|
"http://localhost:4000",
|
|
"sk-key",
|
|
get=lambda *a, **k: _FakeResponse(200, {"data": "nope"}),
|
|
)
|
|
assert isinstance(result, PiSyncError)
|
|
|
|
def test_empty_model_list_is_a_value(self):
|
|
result = fetch_model_ids(
|
|
"http://localhost:4000",
|
|
"sk-key",
|
|
get=lambda *a, **k: _FakeResponse(200, {"data": []}),
|
|
)
|
|
assert isinstance(result, PiSyncError)
|
|
assert "no models" in result.message
|
|
assert result.kind is ListingFailure.EMPTY
|
|
|
|
@pytest.mark.parametrize(
|
|
("get", "kind"),
|
|
[
|
|
(_refused, ListingFailure.UNREACHABLE),
|
|
(lambda *a, **k: _FakeResponse(401), ListingFailure.REJECTED),
|
|
(lambda *a, **k: _FakeResponse(403), ListingFailure.REJECTED),
|
|
(lambda *a, **k: _FakeResponse(500), ListingFailure.OTHER),
|
|
(lambda *a, **k: _FakeResponse(200), ListingFailure.BAD_BODY),
|
|
],
|
|
ids=["unreachable", "401", "403", "500", "bad-body"],
|
|
)
|
|
def test_the_failure_kind_is_decided_where_the_response_is_classified(self, get, kind):
|
|
result = fetch_model_ids("http://localhost:4000", "sk-key", get=get)
|
|
assert isinstance(result, PiSyncError) and result.kind is kind
|
|
|
|
|
|
class TestFetchModelLimits:
|
|
def test_maps_group_limits_and_hits_model_group_info(self):
|
|
captured = {}
|
|
|
|
def fake_get(url, headers, timeout):
|
|
captured["url"] = url
|
|
return _FakeResponse(
|
|
200,
|
|
{
|
|
"data": [
|
|
{"model_group": "m-a", "max_input_tokens": 131072, "max_output_tokens": 8192},
|
|
{"model_group": "m-b", "max_input_tokens": None, "max_output_tokens": None},
|
|
]
|
|
},
|
|
)
|
|
|
|
limits = fetch_model_limits("http://localhost:4000/", "sk-key", get=fake_get)
|
|
assert captured["url"] == "http://localhost:4000/model_group/info"
|
|
assert limits["m-a"] == ModelLimits(context_window=131072, max_tokens=8192)
|
|
assert limits["m-b"] == ModelLimits(context_window=None, max_tokens=None)
|
|
|
|
def test_non_200_degrades_to_no_limits(self):
|
|
assert fetch_model_limits("http://localhost:4000", "sk-key", get=lambda *a, **k: _FakeResponse(403)) == {}
|
|
|
|
def test_network_error_degrades_to_no_limits(self):
|
|
def boom(*a, **k):
|
|
raise requests.ConnectionError("refused")
|
|
|
|
assert fetch_model_limits("http://localhost:4000", "sk-key", get=boom) == {}
|
|
|
|
def test_malformed_body_degrades_to_no_limits(self):
|
|
assert (
|
|
fetch_model_limits(
|
|
"http://localhost:4000", "sk-key", get=lambda *a, **k: _FakeResponse(200, {"data": "nope"})
|
|
)
|
|
== {}
|
|
)
|
|
|
|
|
|
class TestModelsJsonPath:
|
|
def test_env_override_wins(self):
|
|
assert models_json_path({"PI_CODING_AGENT_DIR": "/custom/dir"}) == Path("/custom/dir/models.json")
|
|
|
|
def test_defaults_to_home_pi_agent(self):
|
|
assert models_json_path({}) == Path.home() / ".pi" / "agent" / "models.json"
|
|
|
|
|
|
class TestProviderBlock:
|
|
def test_points_pi_at_proxy_with_env_interpolated_key(self):
|
|
block = provider_block("http://localhost:4000/", ("m-1", "m-2"))
|
|
assert block == {
|
|
"baseUrl": "http://localhost:4000/v1",
|
|
"api": "openai-completions",
|
|
"apiKey": "$LITELLM_PROXY_API_KEY",
|
|
"models": [{"id": "m-1"}, {"id": "m-2"}],
|
|
}
|
|
|
|
def test_known_limits_become_context_window_and_max_tokens(self):
|
|
block = provider_block(
|
|
"http://localhost:4000",
|
|
("m-1", "m-2"),
|
|
{
|
|
"m-1": ModelLimits(context_window=131072, max_tokens=8192),
|
|
"m-2": ModelLimits(context_window=None, max_tokens=None),
|
|
},
|
|
)
|
|
assert block["models"] == [
|
|
{"id": "m-1", "contextWindow": 131072, "maxTokens": 8192},
|
|
{"id": "m-2"},
|
|
]
|
|
|
|
|
|
class TestSyncModelsJson:
|
|
def test_creates_file_and_parent_dirs(self, tmp_path):
|
|
path = tmp_path / "agent" / "models.json"
|
|
assert sync_models_json(path, "http://localhost:4000", ("m-1",)) is None
|
|
written = json.loads(path.read_text())
|
|
assert written["providers"]["litellm"]["baseUrl"] == "http://localhost:4000/v1"
|
|
assert written["providers"]["litellm"]["models"] == [{"id": "m-1"}]
|
|
|
|
def test_preserves_other_providers_and_top_level_keys(self, tmp_path):
|
|
path = tmp_path / "models.json"
|
|
path.write_text(
|
|
json.dumps(
|
|
{
|
|
"somethingElse": True,
|
|
"providers": {
|
|
"ollama": {"baseUrl": "http://localhost:11434/v1"},
|
|
"litellm": {"baseUrl": "http://stale:1234/v1", "models": []},
|
|
},
|
|
}
|
|
)
|
|
)
|
|
assert sync_models_json(path, "http://localhost:4000", ("m-1",)) is None
|
|
written = json.loads(path.read_text())
|
|
assert written["somethingElse"] is True
|
|
assert written["providers"]["ollama"] == {"baseUrl": "http://localhost:11434/v1"}
|
|
assert written["providers"]["litellm"]["baseUrl"] == "http://localhost:4000/v1"
|
|
assert written["providers"]["litellm"]["models"] == [{"id": "m-1"}]
|
|
|
|
def test_write_leaves_no_staging_file_behind(self, tmp_path):
|
|
path = tmp_path / "models.json"
|
|
assert sync_models_json(path, "http://localhost:4000", ("m-1",)) is None
|
|
assert [p.name for p in tmp_path.iterdir()] == ["models.json"]
|
|
|
|
def test_written_file_is_private(self, tmp_path):
|
|
path = tmp_path / "models.json"
|
|
assert sync_models_json(path, "http://localhost:4000", ("m-1",)) is None
|
|
if os.name != "nt":
|
|
assert stat.S_IMODE(path.stat().st_mode) == 0o600
|
|
|
|
path.write_text(json.dumps({"providers": {"other": {"apiKey": "literal-secret"}}}))
|
|
path.chmod(0o644)
|
|
assert sync_models_json(path, "http://localhost:4000", ("m-2",)) is None
|
|
if os.name != "nt":
|
|
assert stat.S_IMODE(path.stat().st_mode) == 0o600
|
|
|
|
def test_concurrent_syncs_do_not_collide(self, tmp_path):
|
|
path = tmp_path / "models.json"
|
|
model_lists = (("m-a",), ("m-b",))
|
|
|
|
def sync(model_ids):
|
|
return sync_models_json(path, "http://localhost:4000", model_ids)
|
|
|
|
with ThreadPoolExecutor(max_workers=2) as executor:
|
|
results = [result for _ in range(30) for result in executor.map(sync, model_lists)]
|
|
|
|
assert results == [None] * 60
|
|
written = json.loads(path.read_text())
|
|
assert written["providers"]["litellm"]["models"] in ([{"id": "m-a"}], [{"id": "m-b"}])
|
|
assert list(tmp_path.glob("models.json.*.tmp")) == []
|
|
|
|
def test_invalid_json_is_a_value_and_file_untouched(self, tmp_path):
|
|
path = tmp_path / "models.json"
|
|
path.write_text("{not json")
|
|
result = sync_models_json(path, "http://localhost:4000", ("m-1",))
|
|
assert isinstance(result, PiSyncError)
|
|
assert path.read_text() == "{not json"
|
|
|
|
def test_non_object_providers_is_a_value(self, tmp_path):
|
|
path = tmp_path / "models.json"
|
|
path.write_text(json.dumps({"providers": ["nope"]}))
|
|
result = sync_models_json(path, "http://localhost:4000", ("m-1",))
|
|
assert isinstance(result, PiSyncError)
|
|
|
|
def test_top_level_non_object_is_a_value(self, tmp_path):
|
|
path = tmp_path / "models.json"
|
|
path.write_text(json.dumps(["nope"]))
|
|
result = sync_models_json(path, "http://localhost:4000", ("m-1",))
|
|
assert isinstance(result, PiSyncError)
|
|
|
|
def test_unwritable_path_is_a_value(self, tmp_path):
|
|
blocker = tmp_path / "agent"
|
|
blocker.write_text("i am a file, not a directory")
|
|
result = sync_models_json(blocker / "models.json", "http://localhost:4000", ("m-1",))
|
|
assert isinstance(result, PiSyncError)
|
|
assert "Could not" in result.message
|