Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
27 changes: 12 additions & 15 deletions src/quant_platform_kit/strategy_lifecycle/ai_provider.py
Original file line number Diff line number Diff line change
Expand Up @@ -143,7 +143,7 @@ def review(self, prompt: str, *, timeout: float = 120.0) -> list["AiCallResult"]
for r in result.results
]

# Fallback: local implementation (no gateway client available)
# A Codex execution job cannot replace independent API reviewers.
return self._review_local(prompt, timeout)

def verify(self, prompt: str, *, timeout: float = 600.0) -> "AiCallResult | None":
Expand Down Expand Up @@ -180,20 +180,17 @@ def _call_single(self, provider: AiProviderConfig, prompt: str, timeout: float)
return self._call_local(provider, prompt, timeout)

def _review_local(self, prompt: str, timeout: float) -> list["AiCallResult"]:
"""Local fallback when gateway client is not installed."""
import concurrent.futures
with concurrent.futures.ThreadPoolExecutor(max_workers=min(len(self.config.reviewers), 3)) as pool:
futures = {pool.submit(self._call_local, c, prompt, timeout): c for c in self.config.reviewers}
results = []
for f in concurrent.futures.as_completed(futures):
try:
results.append(f.result())
except Exception as exc:
results.append(AiCallResult.unavailable(futures[f].label, str(exc)))
return results
"""Report unavailable reviewers when the gateway client is not installed."""
return [
AiCallResult.unavailable(c.label, "ai_gateway_client required for review")
for c in self.config.reviewers
]

def _call_local(self, provider: AiProviderConfig, prompt: str, timeout: float) -> "AiCallResult":
"""Direct HTTP call to AiGateway — used when client library not installed."""
"""Direct Codex execution only when the gateway client is not installed."""
if provider.provider != AiProviderId.CODEX_VPS or provider.task != "execute":
return AiCallResult.unavailable(provider.label, "ai_gateway_client required for this provider/task")

import json as _json
import urllib.error as _urllib_err
import urllib.request as _urllib_req
Expand Down Expand Up @@ -245,10 +242,10 @@ def _call_local(self, provider: AiProviderConfig, prompt: str, timeout: float) -
continue
status = job.get("status")
if status == "succeeded":
return AiCallResult(provider=provider.label, success=True,
return AiCallResult(provider="Codex VPS", success=True,
output=str(job.get("output", "")), raw=job)
if status == "failed":
return AiCallResult(provider=provider.label, success=False,
return AiCallResult(provider="Codex VPS", success=False,
output=job.get("error", "unknown"), raw=job)
return AiCallResult.unavailable(provider.label, "Timeout")
except Exception as exc:
Expand Down
84 changes: 83 additions & 1 deletion tests/test_lifecycle_ai_provider.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,7 +4,9 @@

import json
import unittest
from unittest.mock import patch
from dataclasses import replace
from types import SimpleNamespace
from unittest.mock import Mock, patch

from quant_platform_kit.strategy_lifecycle import ai_provider

Expand All @@ -25,6 +27,86 @@ def read(self) -> bytes:

class AiProviderGatewayFallbackTests(unittest.TestCase):

def test_review_without_sdk_is_unavailable_without_execute_fallback(self) -> None:
reviewers = [ai_provider.AiProviderConfig.claude(), ai_provider.AiProviderConfig.gpt()]
config = ai_provider.AiServiceConfig.safety(reviewers=reviewers)
with patch.object(ai_provider, "_HAS_GATEWAY_CLIENT", False), patch.object(
ai_provider.AiServiceClient, "_call_local",
return_value=ai_provider.AiCallResult(provider="Codex VPS", success=True, output="approve"),
) as local:
results = ai_provider.AiServiceClient(config).review("synthetic review")
self.assertEqual([r.provider for r in results], ["Claude", "GPT"])
self.assertTrue(all(not r.success and not r.output for r in results))
self.assertTrue(all("ai_gateway_client" in r.note for r in results))
local.assert_not_called()

def test_no_sdk_rejects_non_codex_and_non_execute_before_auth_or_http(self) -> None:
codex = ai_provider.AiProviderConfig.codex_vps()
providers = [
ai_provider.AiProviderConfig.claude(), ai_provider.AiProviderConfig.gpt(),
replace(ai_provider.AiProviderConfig.gpt(), task="execute"),
replace(codex, task="analyze"), replace(codex, task="review"),
]
for provider in providers:
for operation in ("_call_single", "verify"):
with self.subTest(provider=provider.provider, task=provider.task, operation=operation):
config = ai_provider.AiServiceConfig.safety(reviewers=[], verifier=provider)
with patch.object(ai_provider, "_HAS_GATEWAY_CLIENT", False), patch.dict(
ai_provider.os.environ, {"CODEX_AUDIT_SERVICE_URL": "https://gateway.example"}, clear=True,
), patch.object(ai_provider, "_fetch_oidc_token", return_value="synthetic") as auth, patch(
"urllib.request.urlopen", return_value=_FakeResponse({"job_id": "synthetic-job"}),
) as http, patch("time.time", side_effect=[0, 1000]):
client = ai_provider.AiServiceClient(config)
result = (client.verify("synthetic", timeout=1) if operation == "verify"
else client._call_single(provider, "synthetic", 1))
self.assertFalse(result.success)
self.assertEqual(result.output, "")
self.assertIn("ai_gateway_client", result.note)
auth.assert_not_called()
http.assert_not_called()

def test_no_sdk_codex_execute_and_verify_preserve_endpoint_identity(self) -> None:
provider = replace(ai_provider.AiProviderConfig.codex_vps(), label="Claude")
for operation in ("execute", "verify"):
with self.subTest(operation=operation):
config = ai_provider.AiServiceConfig(
pattern=ai_provider.AiPattern.RELIABILITY, primary=provider, verifier=provider,
)
with patch.object(ai_provider, "_HAS_GATEWAY_CLIENT", False), patch.dict(
ai_provider.os.environ, {"CODEX_AUDIT_SERVICE_URL": "https://gateway.example"}, clear=True,
), patch.object(ai_provider, "_fetch_oidc_token", return_value="synthetic"), patch(
"urllib.request.urlopen", side_effect=[
_FakeResponse({"job_id": "synthetic-job"}),
_FakeResponse({"status": "succeeded", "output": "synthetic advisory"}),
],
) as http, patch("time.sleep", return_value=None):
result = getattr(ai_provider.AiServiceClient(config), operation)("synthetic", timeout=1)
self.assertTrue(result.success)
self.assertEqual(result.provider, "Codex VPS")
self.assertEqual(result.output, "synthetic advisory")
self.assertEqual(http.call_count, 2)
request = http.call_args_list[0].args[0]
self.assertEqual(request.full_url, "https://gateway.example/v1/ai/execute/jobs")
self.assertEqual(json.loads(request.data)["mode"], "review_only")
self.assertEqual(json.loads(request.data)["task"], "execute")

def test_sdk_analyze_keeps_actual_provider_instead_of_caller_label(self) -> None:
provider = replace(ai_provider.AiProviderConfig.gpt(), label="Claude")
gateway = Mock()
gateway.analyze.return_value = SimpleNamespace(
provider="openai", success=True, output="synthetic advisory", error="",
)
with patch.object(ai_provider, "_HAS_GATEWAY_CLIENT", True), patch.object(
ai_provider, "GatewayConfig", create=True,
), patch.object(ai_provider, "AiGatewayClient", return_value=gateway, create=True):
result = ai_provider.AiServiceClient(
ai_provider.AiServiceConfig.reliability(primary=provider),
).execute("synthetic", timeout=1)
self.assertEqual(result.provider, "openai")
self.assertTrue(result.success)
gateway.analyze.assert_called_once_with("synthetic", model=provider.model, timeout=1)
gateway.execute.assert_not_called()

def test_local_gateway_payload_defaults_to_quant_platform_kit(self) -> None:
requests = []

Expand Down