diff --git a/src/task_scheduler.py b/src/task_scheduler.py index 02fa9970f..beeb4dee5 100644 --- a/src/task_scheduler.py +++ b/src/task_scheduler.py @@ -1969,21 +1969,20 @@ class TaskScheduler: messages.append(datetime_context_msg) messages.append({"role": "user", "content": user_content}) - # Resolve headers from the endpoint's API key + # Recover the registered runtime URL and credentials for this chat route. headers = {} try: - from core.database import SessionLocal, ModelEndpoint - from src.endpoint_resolver import normalize_base, build_headers, same_endpoint_base - from src.auth_helpers import owner_filter + from core.database import SessionLocal + from src.endpoint_resolver import ( + build_chat_url, build_headers, resolve_endpoint_runtime, + resolve_owner_registered_endpoint, + ) db2 = SessionLocal() try: - ep_q = db2.query(ModelEndpoint).filter(ModelEndpoint.is_enabled == True) - ep_q = owner_filter(ep_q, ModelEndpoint, task.owner or None) - eps = ep_q.all() - for ep in eps: - if same_endpoint_base(endpoint_url, ep.base_url): - headers = build_headers(ep.api_key, normalize_base(ep.base_url)) - break + ep = resolve_owner_registered_endpoint(db2, endpoint_url, task.owner or None) + base, api_key = resolve_endpoint_runtime(ep, owner=task.owner or None) + endpoint_url = build_chat_url(base) + headers = build_headers(api_key, base) finally: db2.close() except Exception: @@ -2170,26 +2169,23 @@ class TaskScheduler: endpoint_url, model = self._resolve_defaults(db, task.owner) if not endpoint_url or not model: raise RuntimeError("No model/endpoint configured for research") - endpoint_url = _normalize_chat_endpoint(endpoint_url) # Record the resolved model for the run record (see _execute_task_locked). self._last_run_model = model - # Resolve headers + # Authorize the selected URL before normalization can collapse service paths. try: - from core.database import ModelEndpoint - from src.endpoint_resolver import normalize_base, build_headers, same_endpoint_base - from src.auth_helpers import owner_filter - db2 = db + from src.endpoint_resolver import ( + build_chat_url, build_headers, resolve_endpoint_runtime, + resolve_owner_registered_endpoint, + ) if not headers_from_resolver: - ep_q = db2.query(ModelEndpoint).filter(ModelEndpoint.is_enabled == True) - ep_q = owner_filter(ep_q, ModelEndpoint, task.owner or None) - eps = ep_q.all() - for ep in eps: - if same_endpoint_base(endpoint_url, ep.base_url): - headers = build_headers(ep.api_key, normalize_base(ep.base_url)) - break + ep = resolve_owner_registered_endpoint(db, endpoint_url, task.owner or None) + base, api_key = resolve_endpoint_runtime(ep, owner=task.owner or None) + endpoint_url = build_chat_url(base) + headers = build_headers(api_key, base) except Exception: pass + endpoint_url = _normalize_chat_endpoint(endpoint_url) max_tokens = int(get_setting("research_max_tokens", 8192)) extraction_timeout = int(get_setting("research_extraction_timeout_seconds", 90) or 90) diff --git a/tests/test_aux_llm_owner_scope.py b/tests/test_aux_llm_owner_scope.py index e9bc383bd..f4aeec420 100644 --- a/tests/test_aux_llm_owner_scope.py +++ b/tests/test_aux_llm_owner_scope.py @@ -49,8 +49,9 @@ def test_scheduler_fallbacks_and_research_headers_are_owner_scoped(): assert "owner=task.owner or None" in src assert "headers_from_resolver = False" in src assert "headers_from_resolver = True" in src - assert "from src.auth_helpers import owner_filter" in src - assert "owner_filter(ep_q, ModelEndpoint, task.owner or None)" in src + assert "resolve_owner_registered_endpoint(db2, endpoint_url, task.owner or None)" in src + assert "resolve_owner_registered_endpoint(db, endpoint_url, task.owner or None)" in src + assert src.count("resolve_endpoint_runtime(ep, owner=task.owner or None)") == 2 def test_research_routes_fallbacks_are_owner_scoped(): diff --git a/tests/test_task_scheduler_endpoint_authority.py b/tests/test_task_scheduler_endpoint_authority.py new file mode 100644 index 000000000..4bae69242 --- /dev/null +++ b/tests/test_task_scheduler_endpoint_authority.py @@ -0,0 +1,210 @@ +"""Scheduler credential recovery uses registered endpoint identity after normalization.""" + +from types import SimpleNamespace +from unittest.mock import AsyncMock, MagicMock +from urllib.parse import urlparse + +import pytest + +import core.database as database +import src.endpoint_resolver as resolver +from src.agent_runtime.authority import seal_task_authority +from src.task_scheduler import TaskScheduler +from tests.helpers.database import disposable_database + + +PROVIDERS = [ + ("openai", "https://api.openai.com", "https://api.openai.com/v1/chat/completions"), + ("anthropic_v1", "https://api.anthropic.com/v1", "https://api.anthropic.com/v1/messages"), + ("ollama", "http://192.168.1.5:11434", "http://192.168.1.5:11434/api/chat"), + ("generic", "https://models.example.test/v1", "https://models.example.test/v1/chat/completions"), + ("anthropic_root", "https://api.anthropic.com", "https://api.anthropic.com/v1/messages"), + ("ollama_api", "http://192.168.1.5:11434/api", "http://192.168.1.5:11434/api/chat"), +] + + +@pytest.fixture +def execution(tmp_path, monkeypatch): + with disposable_database(tmp_path) as factory: + monkeypatch.setattr(database, "SessionLocal", factory) + monkeypatch.setattr(resolver, "resolve_url", lambda url: url) + monkeypatch.setattr("src.interactive_gate.wait_for_interactive_quiet", AsyncMock()) + monkeypatch.setattr("src.task_endpoint.resolve_task_candidates", lambda **kwargs: []) + monkeypatch.setattr("src.settings.get_setting", lambda key, default=None: default) + monkeypatch.setattr("src.task_scheduler._resolve_task_timezone", lambda db, task: "UTC") + monkeypatch.setattr("src.tool_index.get_tool_index", lambda: None) + monkeypatch.setattr("src.research_handler.RESEARCH_DATA_DIR", tmp_path / "research") + monkeypatch.setattr("src.event_bus.fire_event", lambda *args: None) + captured = {} + + async def stream(**kwargs): + captured.update(endpoint_url=kwargs["endpoint_url"], headers=kwargs["headers"]) + yield 'data: {"delta": "Complete"}\n\n' + + monkeypatch.setattr("src.agent_loop.stream_agent_loop", stream) + researcher = SimpleNamespace(research=AsyncMock(return_value="Complete"), + get_stats=lambda: {}, findings=[]) + + def research(**kwargs): + captured.update(endpoint_url=kwargs["llm_endpoint"], headers=kwargs["llm_headers"]) + return researcher + + monkeypatch.setattr("src.deep_research.DeepResearcher", research) + + async def run(kind, url, *, execute_agent_task=True, **overrides): + task = SimpleNamespace( + id="task", owner="alice", name="Research task", prompt="Research this topic", + task_type="research" if kind == "research" else "llm", action=None, + endpoint_url=url, model="model", session_id="session", max_steps=1, + crew_member_id=None, character_id=None, + headers={"Authorization": "Bearer caller-secret", "x-api-key": "caller-secret"}, + ) + task.__dict__.update(overrides) + task.request_authority_json = seal_task_authority( + task.prompt, task.task_type, task.action, owner=task.owner) + scheduler = TaskScheduler.__new__(TaskScheduler) + scheduler._session_manager = None + if kind == "agent": + if execute_agent_task: + with factory() as db: + result = await scheduler._execute_llm_task(task, db) + else: + result = await scheduler._run_agent_loop(url, task.model, task, task.session_id) + else: + with factory() as db: + result = await scheduler._execute_research_task(task, db) + assert result == "Complete" + return captured + + yield factory, run + + +@pytest.mark.parametrize("kind", ["agent", "research"]) +@pytest.mark.parametrize("provider, base, chat_url", PROVIDERS) +@pytest.mark.parametrize("selection", ["base", "chat"]) +@pytest.mark.parametrize("server_key", ["server-secret", None]) +async def test_scheduler_recovers_static_credentials_after_chat_normalization(execution, kind, provider, base, chat_url, selection, server_key): + factory, run = execution + with factory() as db: + db.add(database.ModelEndpoint(id="endpoint", name="Endpoint", owner="alice", + base_url=base, api_key=server_key, is_enabled=True)) + db.commit() + captured = await run(kind, base if selection == "base" else chat_url) + assert captured == {"endpoint_url": chat_url, "headers": resolver.build_headers(server_key, base)} + + +@pytest.mark.parametrize("kind", ["agent", "research"]) +@pytest.mark.parametrize("provider, base, chat_url", PROVIDERS) +async def test_scheduler_refreshes_registered_session_credentials(execution, monkeypatch, kind, provider, base, chat_url): + factory, run = execution + with factory() as db: + db.add(database.ModelEndpoint(id="endpoint", name="Endpoint", owner="alice", + base_url=base, api_key="stale-secret", is_enabled=True, + provider_auth_id="server-session")) + db.commit() + runtime = MagicMock(return_value={"base_url": base, "api_key": "refreshed-secret"}) + monkeypatch.setattr("src.chatgpt_subscription.resolve_runtime_credentials", runtime) + captured = await run(kind, base) + runtime.assert_called_once_with("server-session", owner="alice") + assert captured == {"endpoint_url": chat_url, "headers": resolver.build_headers("refreshed-secret", base)} + + +@pytest.mark.parametrize("kind", ["agent", "research"]) +async def test_scheduler_uses_refreshed_server_runtime_url(execution, monkeypatch, kind): + factory, run = execution + base = "https://chatgpt.com/backend-api/codex" + runtime_base = "https://runtime.example.test/v1" + with factory() as db: + db.add(database.ModelEndpoint(id="endpoint", name="Endpoint", owner="alice", + base_url=base, api_key="stale-secret", is_enabled=True, + provider_auth_id="server-session")) + db.commit() + runtime = MagicMock(return_value={"base_url": runtime_base, "api_key": "refreshed-secret"}) + monkeypatch.setattr("src.chatgpt_subscription.resolve_runtime_credentials", runtime) + captured = await run(kind, base) + runtime.assert_called_once_with("server-session", owner="alice") + assert captured == {"endpoint_url": runtime_base + "/chat/completions", + "headers": {"Authorization": "Bearer refreshed-secret"}} + + +@pytest.mark.parametrize("kind", ["agent", "research"]) +async def test_scheduler_refresh_failure_never_uses_static_or_caller_credentials(execution, monkeypatch, kind): + factory, run = execution + base = PROVIDERS[0][1] + with factory() as db: + db.add(database.ModelEndpoint(id="endpoint", name="Endpoint", owner="alice", + base_url=base, api_key="stale-secret", is_enabled=True, + provider_auth_id="server-session")) + db.commit() + runtime = MagicMock(side_effect=ValueError("Session expired")) + monkeypatch.setattr("src.chatgpt_subscription.resolve_runtime_credentials", runtime) + assert (await run(kind, base))["headers"] == {} + runtime.assert_called_once_with("server-session", owner="alice") + + +@pytest.mark.parametrize("kind", ["agent", "research"]) +@pytest.mark.parametrize("provider, base, chat_url", PROVIDERS[:4]) +@pytest.mark.parametrize("state", ["disabled", "foreign_owner", "unregistered", "shared"]) +async def test_scheduler_credentials_require_enabled_visible_registration(execution, monkeypatch, kind, provider, base, chat_url, state): + factory, run = execution + with factory() as db: + if state != "unregistered": + db.add(database.ModelEndpoint(id="endpoint", name="Endpoint", + owner="bob" if state == "foreign_owner" else None if state == "shared" else "alice", + base_url=base, api_key="server-secret", is_enabled=state != "disabled")) + db.commit() + runtime = MagicMock(wraps=resolver.resolve_endpoint_runtime) + monkeypatch.setattr(resolver, "resolve_endpoint_runtime", runtime) + captured = await run(kind, base) + if state == "shared": + assert captured["headers"] == resolver.build_headers("server-secret", base) + assert runtime.call_args.kwargs == {"owner": "alice"} + else: + assert captured["headers"] == {} + runtime.assert_not_called() + + +@pytest.mark.parametrize("kind", ["agent", "research"]) +@pytest.mark.parametrize("provider, base, chat_url", PROVIDERS[:4]) +@pytest.mark.parametrize("change", [ + "scheme", "host", "port", "userinfo", "query", "fragment", + "sibling", "descendant", "models", "responses", +]) +async def test_scheduler_invalid_route_never_recovers_registered_credentials(execution, monkeypatch, kind, provider, base, chat_url, change): + factory, run = execution + with factory() as db: + db.add(database.ModelEndpoint(id="endpoint", name="Endpoint", owner="alice", + base_url=base, api_key="server-secret", is_enabled=True)) + db.commit() + parsed = urlparse(chat_url) + mutations = { + "scheme": parsed._replace(scheme="https" if parsed.scheme == "http" else "http"), + "host": parsed._replace(netloc="attacker.example.test"), + "port": parsed._replace(netloc=f"{parsed.hostname}:4444"), + "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"), + "responses": parsed._replace(path=urlparse(base).path + "/responses"), + } + runtime = MagicMock(wraps=resolver.resolve_endpoint_runtime) + monkeypatch.setattr(resolver, "resolve_endpoint_runtime", runtime) + captured = await run(kind, mutations[change].geturl(), execute_agent_task=False) + assert captured["headers"] == {} + runtime.assert_not_called() + + +async def test_scheduled_research_keeps_configured_resolver_credentials(execution, monkeypatch): + _, run = execution + url = "https://runtime.example.test/v1/chat/completions" + headers = {"Authorization": "Bearer configured-server-secret"} + configured = MagicMock(return_value=(url, "model", headers)) + monkeypatch.setattr(resolver, "resolve_endpoint", configured) + recovery = MagicMock(side_effect=AssertionError("Configured credentials must not be replaced")) + monkeypatch.setattr(resolver, "resolve_owner_registered_endpoint", recovery) + captured = await run("research", None) + assert captured == {"endpoint_url": url, "headers": headers} + assert configured.call_args.kwargs["owner"] == "alice" + recovery.assert_not_called()