diff --git a/src/endpoint_resolver.py b/src/endpoint_resolver.py index e83ee97e0..f4745c459 100644 --- a/src/endpoint_resolver.py +++ b/src/endpoint_resolver.py @@ -281,19 +281,46 @@ def same_endpoint_base(left, right) -> bool: return False +def _registered_endpoint_url_identity(value): + """Compare complete URL paths without collapsing caller-selected suffixes.""" + if not isinstance(value, str): + return None + value = value.strip() + try: + parsed = urlparse(value) + if (parsed.scheme not in {"http", "https"} or not parsed.hostname + or parsed.username is not None or parsed.password is not None + or "?" in value or "#" in value or parsed.params): + return None + port = parsed.port + return (parsed.scheme, parsed.hostname.lower(), + port if port is not None else (443 if parsed.scheme == "https" else 80), + parsed.path.rstrip("/")) + except ValueError: + return None + + def resolve_owner_registered_endpoint(db, endpoint_url: str, owner: Optional[str] = None): """Authorize a caller URL against enabled, owner-visible endpoint rows. + Accept only the registered canonical base or its server-derived chat URL. Request credentials, query strings and fragments are never endpoint identity. Return the server-owned row so runtime credentials come from registration. """ from src.auth_helpers import owner_filter - if not isinstance(endpoint_url, str) or not same_endpoint_base(endpoint_url, endpoint_url): + identity = _registered_endpoint_url_identity(endpoint_url) + if identity is None: raise ValueError("Invalid model endpoint URL") query = db.query(ModelEndpoint).filter(ModelEndpoint.is_enabled.is_(True)) for endpoint in owner_filter(query, ModelEndpoint, owner).all(): - if same_endpoint_base(endpoint_url, endpoint.base_url): + base = normalize_base(endpoint.base_url) + base_identity = _registered_endpoint_url_identity(base) + if base_identity is None: + continue + if identity == base_identity: + return endpoint + if identity == _registered_endpoint_url_identity(build_chat_url(base)): return endpoint raise ValueError("Model endpoint must be enabled and registered for the current owner") diff --git a/tests/test_endpoint_registered_authority.py b/tests/test_endpoint_registered_authority.py index 8af81d748..aa1a224b1 100644 --- a/tests/test_endpoint_registered_authority.py +++ b/tests/test_endpoint_registered_authority.py @@ -2,12 +2,14 @@ from types import SimpleNamespace from unittest.mock import AsyncMock, MagicMock +from urllib.parse import urlparse import pytest from fastapi import HTTPException import core.database as database import routes.assistant_routes as assistant_routes +import routes.model_routes as model_routes import routes.skills_routes as skills_routes import routes.task.task_routes as task_routes import src.endpoint_resolver as resolver @@ -17,15 +19,40 @@ from tests.helpers.database import disposable_database LOCAL = "http://localhost:1234/v1" LAN = "http://192.168.1.20:8000/v1" SHARED = "http://127.0.0.1:11434/api" +PROVIDER_CASES = [ + ("ollama", "http://192.168.1.5:11434", "http://192.168.1.5:11434/api/chat"), + ("openai", "https://api.openai.com", "https://api.openai.com/v1/chat/completions"), + ("anthropic", "https://api.anthropic.com/v1", "https://api.anthropic.com/v1/messages"), + ("local", LOCAL, LOCAL + "/chat/completions"), +] REJECTED = [ "https://unregistered.example/v1", "http://169.254.169.254/latest/meta-data", "http://bob.example/v1", + "http://bob.example/v1/chat/completions", "http://disabled.example/v1", + "http://disabled.example/v1/chat/completions", "http://caller:secret@localhost:1234/v1", + "http://caller@localhost:1234/v1/chat/completions", + "http://:secret@localhost:1234/v1/chat/completions", + "https://localhost:1234/v1", + "http://localhost.example:1234/v1", + "http://localhost:1235/v1", LOCAL + "?api_key=caller-secret", LOCAL + "#fragment", LOCAL + "/other-base", + LOCAL + "/models", + LOCAL + "/completions", + LOCAL + "/responses", + LOCAL + "/chat/completions/descendant", + LOCAL + "/chat/completions?api_key=caller-secret", + LOCAL + "/chat/completions#fragment", + LOCAL + "/chat/completions?", + LOCAL + "/chat/completions#", + SHARED + "/tags", + SHARED + "/generate", + "https://api.openai.com/v1/models", + "https://api.anthropic.com/v1/models", "not-a-url", " ", ] @@ -34,6 +61,7 @@ REJECTED = [ def _request(body=None, owner="alice"): return SimpleNamespace( state=SimpleNamespace(current_user=owner), + app=SimpleNamespace(state=SimpleNamespace(auth_manager=None)), headers={}, json=AsyncMock(return_value=body or {}), ) @@ -46,8 +74,9 @@ def _route(router, method, path): @pytest.fixture def registered_db(tmp_path, monkeypatch): with disposable_database(tmp_path) as factory: - for module in (database, task_routes, assistant_routes, resolver): + for module in (database, task_routes, assistant_routes, model_routes, resolver): monkeypatch.setattr(module, "SessionLocal", factory) + monkeypatch.setattr(resolver, "resolve_url", lambda url: url) with factory() as db: for endpoint_id, owner, url, enabled in [ ("local", "alice", LOCAL + "/", True), @@ -55,10 +84,13 @@ def registered_db(tmp_path, monkeypatch): ("shared", None, SHARED, True), ("bob", "bob", "http://bob.example/v1", True), ("disabled", "alice", "http://disabled.example/v1", False), + *[(endpoint_id, "alice", base, True) + for endpoint_id, base, _ in PROVIDER_CASES if endpoint_id != "local"], ]: db.add(database.ModelEndpoint( id=endpoint_id, name=endpoint_id, owner=owner, base_url=url, is_enabled=enabled, api_key="server-secret", + cached_models='["model"]', pinned_models='["model"]', )) db.add(database.ScheduledTask( id="task", owner="alice", name="Existing task", prompt="Work", @@ -73,6 +105,82 @@ def registered_db(tmp_path, monkeypatch): yield factory +@pytest.mark.parametrize("endpoint_id, base, chat_url", PROVIDER_CASES) +def test_model_catalog_chat_url_resolves_to_registered_endpoint(registered_db, endpoint_id, base, chat_url): + catalog = _route(model_routes.setup_model_routes(MagicMock()), "GET", "/api/models")(_request()) + item = next(item for item in catalog["items"] if item["endpoint_id"] == endpoint_id) + assert item["url"] == chat_url == resolver.build_chat_url(base) + assert item["models"] == ["model"] + with registered_db() as db: + assert resolver.resolve_owner_registered_endpoint_url(db, item["url"], "alice") == base + endpoint = resolver.resolve_owner_registered_endpoint(db, item["url"], "alice") + assert endpoint.id == endpoint_id + assert resolver.resolve_endpoint_runtime(endpoint, owner="alice") == (base, "server-secret") + + +@pytest.mark.parametrize("endpoint_id, base, chat_url", PROVIDER_CASES) +def test_registered_provider_base_remains_valid(registered_db, endpoint_id, base, chat_url): + with registered_db() as db: + assert resolver.resolve_owner_registered_endpoint_url(db, base, "alice") == base + + +@pytest.mark.parametrize("endpoint_id, base, chat_url", PROVIDER_CASES) +@pytest.mark.parametrize("change", [ + "scheme", "host", "port", "zero_port", "userinfo", "query", "fragment", + "sibling", "descendant", "models", "completions", "responses", "nested_chat", +]) +def test_registered_provider_chat_url_rejects_mutations(registered_db, endpoint_id, base, chat_url, change): + parsed = urlparse(chat_url) + mutations = { + "scheme": parsed._replace(scheme="https" if parsed.scheme == "http" else "http"), + "host": parsed._replace(netloc="attacker.example"), + "port": parsed._replace(netloc=f"{parsed.hostname}:{(parsed.port or 443) + 1}"), + "zero_port": parsed._replace(netloc=f"{parsed.hostname}:0"), + "userinfo": parsed._replace(netloc=f"caller:secret@{parsed.netloc}"), + "query": parsed._replace(query="api_key=caller-secret"), + "fragment": parsed._replace(fragment="fragment"), + "sibling": parsed._replace(path=urlparse(base).path + "/sibling"), + "descendant": parsed._replace(path=parsed.path + "/descendant"), + "models": parsed._replace(path=urlparse(base).path + "/models"), + "completions": parsed._replace(path=urlparse(base).path + "/completions"), + "responses": parsed._replace(path=urlparse(base).path + "/responses"), + "nested_chat": parsed._replace(path=parsed.path + "/chat/completions"), + } + with registered_db() as db, pytest.raises(ValueError): + resolver.resolve_owner_registered_endpoint(db, mutations[change].geturl(), "alice") + + +@pytest.mark.parametrize("endpoint_id, base, chat_url", PROVIDER_CASES) +@pytest.mark.parametrize("state", ["disabled", "other_owner"]) +def test_registered_provider_chat_url_requires_enabled_owner_visibility(registered_db, endpoint_id, base, chat_url, state): + with registered_db() as db: + endpoint = db.get(database.ModelEndpoint, endpoint_id) + if state == "disabled": + endpoint.is_enabled = False + else: + endpoint.owner = "bob" + db.commit() + with pytest.raises(ValueError): + resolver.resolve_owner_registered_endpoint(db, chat_url, "alice") + + +@pytest.mark.parametrize("base, chat_url", [ + ("https://api.openai.com/v1", "https://api.openai.com/v1/chat/completions"), + ("https://api.anthropic.com", "https://api.anthropic.com/v1/messages"), + ("https://ollama.com", "https://ollama.com/api/chat"), + ("http://192.168.1.5:11434/v1", "http://192.168.1.5:11434/v1/chat/completions"), +]) +def test_registered_provider_alternate_base_shapes(tmp_path, monkeypatch, base, chat_url): + monkeypatch.setattr(resolver, "resolve_url", lambda url: url) + with disposable_database(tmp_path) as factory, factory() as db: + db.add(database.ModelEndpoint(id="provider", name="Provider", owner="alice", + base_url=base, is_enabled=True, api_key="server-secret")) + db.commit() + assert resolver.build_chat_url(base) == chat_url + for url in (base, chat_url): + assert resolver.resolve_owner_registered_endpoint_url(db, url, "alice") == base + + @pytest.mark.parametrize("url", REJECTED + ["", None, 42]) def test_registered_endpoint_helper_rejects_invalid_or_invisible_url(registered_db, url): with registered_db() as db, pytest.raises(ValueError): @@ -81,7 +189,7 @@ def test_registered_endpoint_helper_rejects_invalid_or_invisible_url(registered_ @pytest.mark.parametrize("url, canonical", [ ("HTTP://LOCALHOST:1234/v1/chat/completions/", LOCAL), - (LAN + "/models", LAN), + (LAN, LAN), (SHARED + "/chat", SHARED), ]) def test_registered_endpoint_helper_preserves_local_lan_and_shared(registered_db, url, canonical): @@ -126,7 +234,7 @@ async def test_task_create_and_update_store_registered_canonical_url(registered_ )) assert result["endpoint_url"] == LOCAL assert (await update(_request(), "task", task_routes.TaskUpdate( - endpoint_url=LOCAL + "/models", + endpoint_url=LOCAL + "/chat/completions", )))["endpoint_url"] == LOCAL with registered_db() as db: created = db.get(database.ScheduledTask, result["id"]) @@ -136,6 +244,25 @@ async def test_task_create_and_update_store_registered_canonical_url(registered_ assert db.get(database.ScheduledTask, "task").endpoint_url == LOCAL +@pytest.mark.asyncio +@pytest.mark.parametrize("endpoint_id, base, chat_url", PROVIDER_CASES) +async def test_task_create_and_edit_accept_catalog_chat_urls(registered_db, endpoint_id, base, chat_url): + catalog = _route(model_routes.setup_model_routes(MagicMock()), "GET", "/api/models")(_request()) + url = next(item["url"] for item in catalog["items"] if item["endpoint_id"] == endpoint_id) + assert url == chat_url + router = task_routes.setup_task_routes(MagicMock()) + create = _route(router, "POST", "/api/tasks") + update = _route(router, "PUT", "/api/tasks/{task_id}") + result = await create(_request(), task_routes.TaskCreate( + name="Accepted", prompt="Work", trigger_type="webhook", endpoint_url=url, + )) + assert result["endpoint_url"] == base + assert (await update(_request(), "task", task_routes.TaskUpdate(endpoint_url=url)))["endpoint_url"] == base + with registered_db() as db: + assert db.get(database.ScheduledTask, result["id"]).endpoint_url == base + assert db.get(database.ScheduledTask, "task").endpoint_url == base + + @pytest.mark.asyncio async def test_task_empty_override_still_restores_default(registered_db): router = task_routes.setup_task_routes(MagicMock()) @@ -177,6 +304,18 @@ async def test_assistant_settings_accept_registered_local_endpoint(registered_db assert db.get(database.CrewMember, "assistant").endpoint_url == LOCAL +@pytest.mark.asyncio +@pytest.mark.parametrize("endpoint_id, base, chat_url", PROVIDER_CASES) +async def test_assistant_switch_accepts_catalog_chat_urls(registered_db, endpoint_id, base, chat_url): + catalog = _route(model_routes.setup_model_routes(MagicMock()), "GET", "/api/models")(_request()) + url = next(item["url"] for item in catalog["items"] if item["endpoint_id"] == endpoint_id) + assert url == chat_url + update = _route(assistant_routes.setup_assistant_routes(MagicMock()), "PATCH", "/api/assistant/settings") + await update(assistant_routes.AssistantSettingsUpdate(endpoint_url=url), _request()) + with registered_db() as db: + assert db.get(database.CrewMember, "assistant").endpoint_url == base + + @pytest.fixture def skill_test_route(registered_db, monkeypatch): manager = SimpleNamespace( @@ -209,24 +348,53 @@ async def test_skill_test_rejects_raw_fallback_before_network_or_execution(skill @pytest.mark.asyncio @pytest.mark.parametrize("server_key", ["server-secret", None]) -async def test_skill_test_uses_registered_runtime_credentials(skill_test_route, registered_db, server_key): +@pytest.mark.parametrize("endpoint_id, base, chat_url", PROVIDER_CASES) +async def test_skill_test_uses_registered_runtime_credentials(skill_test_route, registered_db, server_key, endpoint_id, base, chat_url): import asyncio with registered_db() as db: - db.get(database.ModelEndpoint, "local").api_key = server_key + db.get(database.ModelEndpoint, endpoint_id).api_key = server_key db.commit() test, probe, run = skill_test_route + catalog = _route(model_routes.setup_model_routes(MagicMock()), "GET", "/api/models")(_request()) + url = next(item["url"] for item in catalog["items"] if item["endpoint_id"] == endpoint_id) + assert url == chat_url await test(_request({ - "endpoint_url": "HTTP://LOCALHOST:1234/v1/chat/completions/", "model": "model", + "endpoint_url": url, "model": "model", "api_key": "attacker", "headers": {"Authorization": "Bearer attacker", "x-api-key": "attacker", "Host": "169.254.169.254"}, }), "skill") await asyncio.sleep(0) headers = {"Authorization": "Bearer server-secret"} if server_key else {} - probe.assert_called_once_with(LOCAL + "/chat/completions", headers=headers) - assert run.await_args.args[4:8] == (LOCAL + "/chat/completions", "model", headers, "alice") + if endpoint_id == "anthropic": + headers = {"anthropic-version": "2023-06-01"} + if server_key: + headers["x-api-key"] = server_key + probe.assert_called_once_with(chat_url, headers=headers) + assert run.await_args.args[4:8] == (chat_url, "model", headers, "alice") assert skills_routes._skill_test_jobs[("alice", "skill")]["_run"]["headers"] == headers +@pytest.mark.asyncio +async def test_skill_test_uses_server_owned_session_credentials(skill_test_route, registered_db, monkeypatch): + import asyncio + + with registered_db() as db: + db.get(database.ModelEndpoint, "local").provider_auth_id = "server-session" + db.commit() + runtime = MagicMock(return_value={"base_url": LAN, "api_key": "server-runtime-secret"}) + monkeypatch.setattr("src.chatgpt_subscription.resolve_runtime_credentials", runtime) + test, probe, run = skill_test_route + await test(_request({ + "endpoint_url": LOCAL + "/chat/completions", "model": "model", + "api_key": "attacker", "headers": {"Authorization": "Bearer attacker"}, + }), "skill") + await asyncio.sleep(0) + runtime.assert_called_once_with("server-session", owner="alice") + headers = {"Authorization": "Bearer server-runtime-secret"} + probe.assert_called_once_with(LAN + "/chat/completions", headers=headers) + assert run.await_args.args[4:8] == (LAN + "/chat/completions", "model", headers, "alice") + + @pytest.mark.asyncio async def test_skill_test_prefers_utility_and_ignores_request_headers(skill_test_route, monkeypatch): import asyncio