fix(endpoints): accept registered canonical chat urls

This commit is contained in:
Alexandre Teixeira 2026-10-06 18:04:59 +01:00
parent 82cfd7d69a
commit 622738bda7
2 changed files with 205 additions and 10 deletions

View file

@ -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")

View file

@ -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