mirror of
https://github.com/pewdiepie-archdaemon/odysseus.git
synced 2026-10-07 14:37:55 +00:00
fix(endpoints): accept registered canonical chat urls
This commit is contained in:
parent
82cfd7d69a
commit
622738bda7
2 changed files with 205 additions and 10 deletions
|
|
@ -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")
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue