test_model_service.py 8.0 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238
  1. from pathlib import Path
  2. from tests.conftest import build_fastapi_test_client, prepare_known_service_import
  3. def test_model_service_post_contract_supports_models_and_providers(
  4. tmp_path: Path,
  5. monkeypatch,
  6. ) -> None:
  7. prepare_known_service_import("model-gateway-service")
  8. from app.bootstrap.app import create_app
  9. from app.db.models import Base
  10. from core_db import create_session_factory
  11. from sqlalchemy import create_engine
  12. database_url = f"sqlite:///{tmp_path / 'models.db'}"
  13. monkeypatch.setenv("AGENT_PLATFORM_DATABASE_URL", database_url)
  14. engine = create_engine(database_url, connect_args={"check_same_thread": False})
  15. Base.metadata.create_all(engine)
  16. app = create_app()
  17. app.state.session_factory = create_session_factory(engine)
  18. client = build_fastapi_test_client(app)
  19. provider_response = client.post(
  20. "/models/providers/create",
  21. json={
  22. "name": "Local OpenAI Compatible",
  23. "providerType": "openai_compatible",
  24. "baseUrl": "http://127.0.0.1:11434/v1",
  25. "apiKey": "local-secret",
  26. "models": [
  27. {
  28. "modelId": "llama3.1",
  29. "displayName": "Llama 3.1",
  30. "modelType": "chat",
  31. }
  32. ],
  33. "defaultModel": "llama3.1",
  34. },
  35. )
  36. assert provider_response.status_code == 200
  37. provider_payload = provider_response.json()["data"]
  38. assert provider_payload["apiKeyRef"] == "loc***masked"
  39. assert provider_payload["models"][0]["modelId"] == "llama3.1"
  40. providers_response = client.post(
  41. "/models/providers/list",
  42. json={"page": 1, "pageSize": 20},
  43. )
  44. assert providers_response.status_code == 200
  45. assert providers_response.json()["data"]["total"] == 1
  46. discover_response = client.post(
  47. "/models/providers/discover",
  48. json={"providerId": provider_payload["id"]},
  49. )
  50. assert discover_response.status_code == 200
  51. assert discover_response.json()["data"]["models"][0]["modelId"] == "llama3.1"
  52. model_response = client.post(
  53. "/models/create",
  54. json={
  55. "name": "Local Chat",
  56. "providerId": provider_payload["id"],
  57. "providerType": "openai_compatible",
  58. "modelName": "llama3.1",
  59. "capabilities": ["chat"],
  60. "timeoutSeconds": 30,
  61. },
  62. )
  63. assert model_response.status_code == 200
  64. model_payload = model_response.json()["data"]
  65. assert model_payload["modelName"] == "llama3.1"
  66. assert model_payload["providerId"] == provider_payload["id"]
  67. assert model_payload["providerBaseUrl"] == "http://127.0.0.1:11434/v1"
  68. assert model_payload["hasProviderApiKey"] is True
  69. assert "code" not in model_payload
  70. models_response = client.post(
  71. "/models/list",
  72. json={"page": 1, "pageSize": 20, "keyword": "local"},
  73. )
  74. assert models_response.status_code == 200
  75. assert models_response.json()["data"]["total"] == 1
  76. update_response = client.post(
  77. "/models/update",
  78. json={
  79. "modelId": model_payload["id"],
  80. "name": "Local Chat Updated",
  81. "defaultTemperature": 0.2,
  82. },
  83. )
  84. assert update_response.status_code == 200
  85. assert update_response.json()["data"]["defaultTemperature"] == 0.2
  86. delete_response = client.post(
  87. "/models/delete",
  88. json={"modelId": model_payload["id"]},
  89. )
  90. assert delete_response.status_code == 200
  91. assert delete_response.json()["data"]["deleted"] is True
  92. def test_model_provider_client_supports_anthropic_messages(monkeypatch) -> None:
  93. prepare_known_service_import("model-gateway-service")
  94. import app.infrastructure.provider as provider_module
  95. from app.bootstrap.settings import ModelGatewayServiceSettings
  96. from app.infrastructure.provider import ModelProviderClient
  97. from core_domain import ChatCompletionRequestContract
  98. captured: dict[str, object] = {}
  99. class FakeResponse:
  100. text = "{}"
  101. def raise_for_status(self) -> None:
  102. return None
  103. def json(self) -> dict[str, object]:
  104. return {
  105. "model": "claude-3-5-sonnet-20241022",
  106. "content": [{"type": "text", "text": "ready"}],
  107. "stop_reason": "end_turn",
  108. "usage": {"input_tokens": 12, "output_tokens": 3},
  109. }
  110. class FakeClient:
  111. def __init__(self, *, timeout: float) -> None:
  112. captured["timeout"] = timeout
  113. def __enter__(self) -> "FakeClient":
  114. return self
  115. def __exit__(self, exc_type: object, exc: object, tb: object) -> None:
  116. return None
  117. def post(
  118. self,
  119. url: str,
  120. *,
  121. json: dict[str, object],
  122. headers: dict[str, str]) -> FakeResponse:
  123. captured["url"] = url
  124. captured["json"] = json
  125. captured["headers"] = headers
  126. return FakeResponse()
  127. monkeypatch.setattr(provider_module.httpx, "Client", FakeClient)
  128. client = ModelProviderClient(settings=ModelGatewayServiceSettings())
  129. response = client.create_chat_completion(
  130. ChatCompletionRequestContract(
  131. model="claude-3-5-sonnet-20241022",
  132. messages=[
  133. {"role": "system", "content": "Be concise."},
  134. {"role": "user", "content": "Ping"},
  135. ],
  136. max_tokens=128),
  137. provider_type="anthropic",
  138. provider_base_url="https://api.anthropic.com",
  139. provider_api_key="sk-test",
  140. timeout_seconds=15)
  141. assert captured["url"] == "https://api.anthropic.com/v1/messages"
  142. assert captured["headers"] == {
  143. "content-type": "application/json",
  144. "x-api-key": "sk-test",
  145. "anthropic-version": "2023-06-01",
  146. }
  147. assert captured["json"] == {
  148. "model": "claude-3-5-sonnet-20241022",
  149. "max_tokens": 128,
  150. "messages": [{"role": "user", "content": "Ping"}],
  151. "system": "Be concise.",
  152. }
  153. assert response.content == "ready"
  154. assert response.finish_reason == "end_turn"
  155. def test_model_service_backfills_legacy_model_connections_as_providers(
  156. tmp_path: Path,
  157. monkeypatch,
  158. ) -> None:
  159. prepare_known_service_import("model-gateway-service")
  160. from app.bootstrap.app import create_app
  161. from app.db.models import Base
  162. from core_db import create_session_factory
  163. from sqlalchemy import create_engine
  164. database_url = f"sqlite:///{tmp_path / 'models.db'}"
  165. monkeypatch.setenv("AGENT_PLATFORM_DATABASE_URL", database_url)
  166. engine = create_engine(database_url, connect_args={"check_same_thread": False})
  167. Base.metadata.create_all(engine)
  168. app = create_app()
  169. app.state.session_factory = create_session_factory(engine)
  170. client = build_fastapi_test_client(app)
  171. model_response = client.post(
  172. "/models/create",
  173. json={
  174. "name": "Legacy Anthropic",
  175. "providerType": "anthropic",
  176. "providerBaseUrl": "https://api.anthropic.com",
  177. "providerApiKey": "sk-legacy",
  178. "modelName": "claude-3-5-sonnet-20241022",
  179. "capabilities": ["chat"],
  180. },
  181. )
  182. assert model_response.status_code == 200
  183. assert model_response.json()["data"]["providerId"] is None
  184. providers_response = client.post(
  185. "/models/providers/list",
  186. json={"page": 1, "pageSize": 20},
  187. )
  188. assert providers_response.status_code == 200
  189. providers_payload = providers_response.json()["data"]
  190. assert providers_payload["total"] == 1
  191. provider_payload = providers_payload["items"][0]
  192. assert provider_payload["providerType"] == "anthropic"
  193. assert provider_payload["baseUrl"] == "https://api.anthropic.com"
  194. assert provider_payload["models"][0]["modelId"] == "claude-3-5-sonnet-20241022"
  195. models_response = client.post(
  196. "/models/list",
  197. json={"page": 1, "pageSize": 20},
  198. )
  199. assert models_response.status_code == 200
  200. assert models_response.json()["data"]["items"][0]["providerId"] == provider_payload["id"]