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
15 changes: 10 additions & 5 deletions e2e/test_evaluator_plugin.py
Original file line number Diff line number Diff line change
Expand Up @@ -49,7 +49,10 @@
from nemo_platform import APIConnectionError, APIStatusError, NeMoPlatform
from nemo_platform.types.inference import ModelProvider
from nemo_platform_plugin.client.adapter import client_from_platform
from nemo_platform_plugin.inference_middleware import BackendFormat
from nemo_platform_plugin.jobs.client import JobsClient
from nemo_platform_plugin.models.client import ModelsClient
from nemo_platform_plugin.models.types import CreateModelEntityRequest
from nmp.testing import add_mock_provider, short_unique_name, wait_for_model_entity
from nmp.testing.e2e import wait_for_platform_job
from nmp.testing.utils import ensure_passthrough_virtual_model
Expand Down Expand Up @@ -263,13 +266,15 @@ def _create_ready_mock_model(
mock_response_body=mock_response_body,
should_autoprovision_virtual_model=False,
)
sdk.models.create(
client_from_platform(sdk, ModelsClient).create_model(
workspace=workspace,
name=name,
backend_format="OPENAI_CHAT",
model_providers=[f"{workspace}/{provider.name}"],
body=CreateModelEntityRequest(
name=name,
backend_format=BackendFormat.OPENAI_CHAT,
model_providers=[f"{workspace}/{provider.name}"],
),
exist_ok=True,
)
).data()
wait_for_model_entity(
sdk,
workspace,
Expand Down
4 changes: 2 additions & 2 deletions packages/nemo_nb/tests/test_myst_stripping.py
Original file line number Diff line number Diff line change
Expand Up @@ -246,7 +246,7 @@ def test_mixed_content():

```python
# Regular code
sdk.models.deploy()
sdk.models.create_deployment()
```

:::{warning}
Expand All @@ -267,4 +267,4 @@ def test_mixed_content():
assert "Use the CLI for faster deployment." in result
assert "Make sure GPU resources are configured." in result
assert "```python" in result
assert "sdk.models.deploy()" in result
assert "sdk.models.create_deployment()" in result
6 changes: 5 additions & 1 deletion packages/nemo_nb/tests/test_notebook_splitting.py
Original file line number Diff line number Diff line change
Expand Up @@ -399,7 +399,11 @@ def test_real_world_tab_set_example(self):
{"cell_type": "code", "metadata": {"language": "bash"}, "source": ["nemo models list\n"]},
{"cell_type": "markdown", "source": [":::\n"]},
{"cell_type": "markdown", "source": [":::{tab-item} Python SDK\n", ":sync: python-sdk\n"]},
{"cell_type": "code", "metadata": {"language": "python"}, "source": ["client.models.list()\n"]},
{
"cell_type": "code",
"metadata": {"language": "python"},
"source": ["client_from_platform(client, ModelsClient).list_models()\n"],
},
{"cell_type": "markdown", "source": [":::\n"]},
{"cell_type": "markdown", "source": ["::::\n"]},
]
Expand Down
10 changes: 5 additions & 5 deletions packages/nemo_nb/tests/test_strip_type_checker_comments.py
Original file line number Diff line number Diff line change
Expand Up @@ -16,9 +16,9 @@ def test_strip_ty_ignore_comments():
"metadata": {"language": "python"},
"source": [
"# This is a regular comment\n",
"sdk.models.get_openai_route_base_url()\n",
"sdk.models.get_model_entity_route_openai_url(entity) # ty: ignore[unresolved-reference]\n",
"sdk.models.get_provider_route_openai_url(provider) # ty: ignore[unresolved-reference]\n",
"client_from_platform(sdk, ModelsClient).get_openai_route_base_url()\n",
"client_from_platform(sdk, ModelsClient).get_model_entity_route_openai_url(entity) # ty: ignore[unresolved-reference]\n",
"client_from_platform(sdk, ModelsClient).get_provider_route_openai_url(provider) # ty: ignore[unresolved-reference]\n",
],
}
]
Expand All @@ -35,8 +35,8 @@ def test_strip_ty_ignore_comments():
assert "# This is a regular comment" in result

# Verify that the code lines are still there (without the ty: comments)
assert "sdk.models.get_model_entity_route_openai_url(entity)" in result
assert "sdk.models.get_provider_route_openai_url(provider)" in result
assert "client_from_platform(sdk, ModelsClient).get_model_entity_route_openai_url(entity)" in result
assert "client_from_platform(sdk, ModelsClient).get_provider_route_openai_url(provider)" in result


def test_strip_type_ignore_comments():
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -16,8 +16,10 @@

from nemo_platform import AsyncNeMoPlatform
from nemo_platform.types.inference import ModelProvider
from nemo_platform.types.models import ModelEntity
from nemo_platform_ext.config import get_context
from nemo_platform_plugin.client.adapter import client_from_platform
from nemo_platform_plugin.models.client import AsyncModelsClient
from nemo_platform_plugin.models.types import ModelEntity
from nooa.unifiedllm import CompletionClient, UnifiedLLM

_PLACEHOLDER_API_KEY = "not-needed"
Expand Down Expand Up @@ -90,8 +92,9 @@ def _completion_client(
served_model_name: str,
) -> CompletionClient:
"""Build a Nooa client that routes through the Model Entity's Platform URL."""
api_base = client.models.get_model_entity_route_openai_url(model_entity)
extra_headers = dict(client.models.get_client_default_headers())
models = client_from_platform(client, AsyncModelsClient)
api_base = models.get_model_entity_route_openai_url(model_entity)
extra_headers = dict(models.default_headers)
# Inference Gateway's direct passthrough session preserves compressed
# response bytes. Nooa needs decoded JSON/SSE, so make that requirement
# explicit at this adapter boundary for every configured agent client.
Expand Down Expand Up @@ -161,6 +164,7 @@ async def resolve_model_clients(
refs: ConfiguredModelRefs | None = None,
) -> ConfiguredModelClients:
"""Resolve configured Model Entities and construct each distinct client once."""
models = client_from_platform(client, AsyncModelsClient)
selected = refs or configured_model_refs()
resolved: dict[str, UnifiedLLM] = {}
provider_cache: dict[str, ModelProvider] = {}
Expand All @@ -169,7 +173,7 @@ async def resolve_model_clients(
if model_ref in resolved:
continue
workspace, name = _parse_model_ref(model_ref)
entity = await client.models.retrieve(name, workspace=workspace)
entity = (await models.get_model(name=name, workspace=workspace)).data()
served_model_name = await _served_model_name(client, entity, provider_cache)
resolved[model_ref] = _completion_client(client, entity, served_model_name)
except Exception as resolution_error:
Expand Down
178 changes: 109 additions & 69 deletions packages/nemo_platform_plugin/tests/test_nooa_model_client.py
Original file line number Diff line number Diff line change
Expand Up @@ -2,7 +2,7 @@
# SPDX-License-Identifier: Apache-2.0

from types import SimpleNamespace
from unittest.mock import AsyncMock, MagicMock
from unittest.mock import AsyncMock, MagicMock, patch

import pytest
from nemo_platform_plugin import nooa_model_client
Expand All @@ -18,6 +18,46 @@
)


def _model_entity(**attrs):
defaults = dict(
workspace="default",
name="model",
backend_format="OPENAI_CHAT",
model_providers=[],
api_endpoint=None,
)
defaults.update(attrs)
return SimpleNamespace(**defaults)


def _mock_models_client(models):
"""Build a mock ``AsyncModelsClient`` returned by ``client_from_platform``.

``models`` maps ``(workspace, name)`` to the Model Entity ``get_model`` should
return. ``get_model`` is an AsyncMock; its awaited result must expose a
``.data()`` accessor mirroring ``NemoResponse``.
"""
client = MagicMock()
client.default_headers = {}
client.get_model_entity_route_openai_url.return_value = "http://platform/model/example/-/v1"

async def _get_model(**kwargs):
entity = models.get((kwargs.get("workspace"), kwargs.get("name")))
return SimpleNamespace(data=lambda: entity)

client.get_model = AsyncMock(side_effect=_get_model)
return client


def _patch_models_client(models):
client = _mock_models_client(models)
ctx = patch(
"nemo_platform_plugin.nooa_model_client.client_from_platform",
return_value=client,
)
return ctx, client


def test_configured_model_refs_uses_default_for_missing_fast(monkeypatch):
monkeypatch.setattr(
nooa_model_client,
Expand All @@ -43,30 +83,28 @@ def test_configured_model_refs_requires_default(monkeypatch):


async def test_resolve_model_clients_deduplicates_same_model(monkeypatch):
model_entity = SimpleNamespace(
model_entity = _model_entity(
workspace="default",
name="gpt-4-1",
backend_format="OPENAI_CHAT",
model_providers=[],
api_endpoint=None,
)
client = MagicMock()
client.models.retrieve = AsyncMock(return_value=model_entity)
client.models.get_model_entity_route_openai_url.return_value = "http://platform/model/gpt-4-1/-/v1"
default_headers = {"x-test": "value"}
client.models.get_client_default_headers.return_value = default_headers
models = {("default", "gpt-4-1"): model_entity}
ctx, client = _patch_models_client(models)
client.get_model_entity_route_openai_url.return_value = "http://platform/model/gpt-4-1/-/v1"
client.default_headers = {"x-test": "value"}
completion_client = MagicMock()
factory = MagicMock(return_value=completion_client)
monkeypatch.setattr(nooa_model_client, "CompletionClient", factory)
with ctx:
monkeypatch.setattr(nooa_model_client, "CompletionClient", factory)

result = await resolve_model_clients(
client,
ConfiguredModelRefs(default="default/gpt-4-1", fast="default/gpt-4-1"),
)
result = await resolve_model_clients(
MagicMock(),
ConfiguredModelRefs(default="default/gpt-4-1", fast="default/gpt-4-1"),
)

assert result.default is completion_client
assert result.fast is completion_client
client.models.retrieve.assert_awaited_once_with("gpt-4-1", workspace="default")
client.get_model.assert_awaited_once_with(name="gpt-4-1", workspace="default")
factory.assert_called_once_with(
"openai/gpt-4-1",
api_base="http://platform/model/gpt-4-1/-/v1",
Expand All @@ -76,28 +114,26 @@ async def test_resolve_model_clients_deduplicates_same_model(monkeypatch):
drop_params=True,
_skip_responses_api_bridge=True,
)
assert default_headers == {"x-test": "value"}
assert client.default_headers == {"x-test": "value"}


async def test_resolve_model_clients_uses_anthropic_route_shape(monkeypatch):
model_entity = SimpleNamespace(
model_entity = _model_entity(
workspace="default",
name="claude-sonnet-4",
backend_format="ANTHROPIC_MESSAGES",
model_providers=[],
api_endpoint=None,
)
client = MagicMock()
client.models.retrieve = AsyncMock(return_value=model_entity)
client.models.get_model_entity_route_openai_url.return_value = "http://platform/model/claude-sonnet-4/-/v1"
client.models.get_client_default_headers.return_value = {}
models = {("default", "claude-sonnet-4"): model_entity}
ctx, client = _patch_models_client(models)
client.get_model_entity_route_openai_url.return_value = "http://platform/model/claude-sonnet-4/-/v1"
factory = MagicMock(return_value=MagicMock())
monkeypatch.setattr(nooa_model_client, "CompletionClient", factory)
with ctx:
monkeypatch.setattr(nooa_model_client, "CompletionClient", factory)

await resolve_model_clients(
client,
ConfiguredModelRefs(default="default/claude-sonnet-4", fast="default/claude-sonnet-4"),
)
await resolve_model_clients(
MagicMock(),
ConfiguredModelRefs(default="default/claude-sonnet-4", fast="default/claude-sonnet-4"),
)

factory.assert_called_once_with(
"anthropic/claude-sonnet-4",
Expand All @@ -110,35 +146,34 @@ async def test_resolve_model_clients_uses_anthropic_route_shape(monkeypatch):


async def test_resolve_model_clients_rejects_unsupported_backend_format():
model_entity = SimpleNamespace(
model_entity = _model_entity(
workspace="default",
name="responses-only",
backend_format="OPENAI_RESPONSES",
model_providers=[],
api_endpoint=None,
)
client = MagicMock()
client.models.retrieve = AsyncMock(return_value=model_entity)

with pytest.raises(ValueError, match="unsupported backend format 'OPENAI_RESPONSES'"):
await resolve_model_clients(
client,
ConfiguredModelRefs(default="default/responses-only", fast="default/responses-only"),
)

models = {("default", "responses-only"): model_entity}
ctx, _client = _patch_models_client(models)

with ctx:
with pytest.raises(ValueError, match="unsupported backend format 'OPENAI_RESPONSES'"):
await resolve_model_clients(
MagicMock(),
ConfiguredModelRefs(default="default/responses-only", fast="default/responses-only"),
)

async def test_resolve_model_clients_requires_workspace_qualified_refs():
client = MagicMock()

with pytest.raises(ValueError, match="workspace/name"):
await resolve_model_clients(
client,
ConfiguredModelRefs(default="unqualified", fast="unqualified"),
)
async def test_resolve_model_clients_requires_workspace_qualified_refs(monkeypatch):
ctx, _client = _patch_models_client({})
with ctx:
with pytest.raises(ValueError, match="workspace/name"):
await resolve_model_clients(
MagicMock(),
ConfiguredModelRefs(default="unqualified", fast="unqualified"),
)


async def test_resolve_model_clients_uses_provider_served_name(monkeypatch):
model_entity = SimpleNamespace(
model_entity = _model_entity(
workspace="default",
name="gpt-5-6-sol",
backend_format="OPENAI_CHAT",
Expand All @@ -153,20 +188,19 @@ async def test_resolve_model_clients_uses_provider_served_name(monkeypatch):
)
]
)
client = MagicMock()
client.models.retrieve = AsyncMock(return_value=model_entity)
client.inference.providers.retrieve = AsyncMock(return_value=provider)
client.models.get_model_entity_route_openai_url.return_value = "http://platform/model/gpt-5-6-sol/-/v1"
client.models.get_client_default_headers.return_value = {}
ctx, client = _patch_models_client({("default", "gpt-5-6-sol"): model_entity})
client.get_model_entity_route_openai_url.return_value = "http://platform/model/gpt-5-6-sol/-/v1"
factory = MagicMock(return_value=MagicMock())
monkeypatch.setattr(nooa_model_client, "CompletionClient", factory)
sdk = MagicMock()
sdk.inference.providers.retrieve = AsyncMock(return_value=provider)
with ctx:
monkeypatch.setattr(nooa_model_client, "CompletionClient", factory)

await resolve_model_clients(
client,
ConfiguredModelRefs(default="default/gpt-5-6-sol", fast="default/gpt-5-6-sol"),
)
await resolve_model_clients(
sdk,
ConfiguredModelRefs(default="default/gpt-5-6-sol", fast="default/gpt-5-6-sol"),
)

client.inference.providers.retrieve.assert_awaited_once_with("openai", workspace="default")
factory.assert_called_once_with(
"openai/gpt-5.6-sol",
api_base="http://platform/model/gpt-5-6-sol/-/v1",
Expand Down Expand Up @@ -228,23 +262,29 @@ async def test_model_clients_close_fast_after_default_close_fails():


async def test_resolve_model_clients_closes_constructed_client_after_failure(monkeypatch):
default_entity = SimpleNamespace(
default_entity = _model_entity(
workspace="default",
name="quality",
backend_format="OPENAI_CHAT",
model_providers=[],
api_endpoint=SimpleNamespace(model_id="quality"),
)
client = MagicMock()
client.models.retrieve = AsyncMock(side_effect=[default_entity, RuntimeError("fast resolution failed")])

async def _flaky_get_model(**kwargs):
if kwargs.get("name") == "quality":
return SimpleNamespace(data=lambda: default_entity)
raise RuntimeError("fast resolution failed")

ctx, client = _patch_models_client({})
client.get_model = AsyncMock(side_effect=_flaky_get_model)
constructed = MagicMock()
constructed.aclose = AsyncMock()
monkeypatch.setattr(nooa_model_client, "_completion_client", MagicMock(return_value=constructed))
with ctx:
monkeypatch.setattr(nooa_model_client, "_completion_client", MagicMock(return_value=constructed))

with pytest.raises(RuntimeError, match="fast resolution failed"):
await resolve_model_clients(
client,
ConfiguredModelRefs(default="default/quality", fast="default/fast"),
)
with pytest.raises(RuntimeError, match="fast resolution failed"):
await resolve_model_clients(
MagicMock(),
ConfiguredModelRefs(default="default/quality", fast="default/fast"),
)

constructed.aclose.assert_awaited_once()
Loading
Loading