litellm/tests/test_litellm/proxy/client/cli/test_pi.py
Devin AI 19c43eb875 test(cli): drop structural StrEnum source check; smoke job covers the 3.10 import
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
2026-09-14 07:08:24 +00:00

301 lines
12 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,
)
def test_listing_failure_is_str_enum():
assert issubclass(ListingFailure, str)
assert ListingFailure.REJECTED.value == "rejected"
assert ListingFailure("rejected") is ListingFailure.REJECTED
assert str(ListingFailure.REJECTED.value) == "rejected"
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