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
47 changes: 47 additions & 0 deletions aqueduct/gateway/tests/test_endpoints.py
Original file line number Diff line number Diff line change
Expand Up @@ -1219,6 +1219,53 @@ def test_list_models(self):
req = requests[0]
self.assertIn("models", req.path, "Request endpoint should be for model listing.")

def test_model_group_info_exposes_context_window(self):
"""The LiteLLM /model_group/info endpoint exposes token limits, so
LiteLLM-aware clients read the context window from the gateway."""
from importlib import import_module

models_view = import_module("gateway.views.models")

config = {
"model_list": [
{
"model_name": self.model,
"litellm_params": {"model": f"openai/{self.model}"},
"model_info": {
"id": self.model,
"max_tokens": 262144,
"max_output_tokens": 32768,
"supports_vision": True,
},
},
{
"model_name": "explicit-model",
"litellm_params": {"model": "openai/explicit-model"},
"model_info": {"max_input_tokens": 999999, "max_tokens": 262144},
},
{
"model_name": "no-info-model",
"litellm_params": {"model": "openai/no-info-model"},
},
]
}
with patch.object(models_view, "get_router_config", return_value=config):
response = self.client.get(
"/model_group/info", content_type="application/json", headers=self.headers
)

self.assertEqual(response.status_code, 200)
entries = {entry["model_group"]: entry for entry in response.json()}
info = entries[self.model]["model_info"]
# max_input_tokens is derived from max_tokens (the context length)
self.assertEqual(info["max_input_tokens"], 262144)
self.assertEqual(info["max_output_tokens"], 32768)
self.assertTrue(info["supports_vision"])
# An explicit max_input_tokens in the config is never overwritten
self.assertEqual(entries["explicit-model"]["model_info"]["max_input_tokens"], 999999)
# Models without token limits get no token fields at all
self.assertNotIn("max_input_tokens", entries["no-info-model"]["model_info"])

def test_list_models_with_invalid_token(self):
"""
Sends a request to list available models from the vLLM server with an invalid API key.
Expand Down
2 changes: 2 additions & 0 deletions aqueduct/gateway/urls.py
Original file line number Diff line number Diff line change
Expand Up @@ -21,6 +21,8 @@
# Models endpoints
path("models", views.models, name="models"),
path("v1/models", views.models, name="v1_models"),
# LiteLLM-style rich model metadata endpoint (probed by LiteLLM-aware clients)
path("model_group/info", views.model_group_info, name="model_group_info"),
# Speech endpoint
path("audio/speech", views.speech, name="speech"),
path("v1/audio/speech", views.speech, name="v1_speech"),
Expand Down
3 changes: 2 additions & 1 deletion aqueduct/gateway/views/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,7 +4,7 @@
from .embeddings import embeddings
from .files import file, file_content, files
from .image_generation import image_generation
from .models import models
from .models import model_group_info, models
from .responses import create_response, get_response_input_items, response
from .speech import speech
from .transcriptions import transcriptions
Expand All @@ -30,6 +30,7 @@
"files",
"get_response_input_items",
"image_generation",
"model_group_info",
"models",
"response",
"speech",
Expand Down
50 changes: 50 additions & 0 deletions aqueduct/gateway/views/models.py
Original file line number Diff line number Diff line change
Expand Up @@ -15,6 +15,27 @@
MODEL_CREATION_TIMESTAMP = int(timezone.now().timestamp())


def _model_group_info_entry(model: dict[str, Any]) -> dict[str, Any]:
"""Build a LiteLLM ``/model_group/info`` entry for a router config model.

The generated router config only carries ``max_tokens`` (the context
length); LiteLLM-aware clients read ``max_input_tokens`` as the context
window, so derive it here — in this HTTP response only, never in the
config itself (a config-level value would re-trigger the Router's
tiktoken pre-checks).
"""
model_info = dict(model.get("model_info") or {})
if "max_input_tokens" not in model_info:
max_tokens = model_info.get("max_tokens")
if isinstance(max_tokens, int) and not isinstance(max_tokens, bool):
model_info["max_input_tokens"] = max_tokens
return {
"model_group": model["model_name"],
"model_name": model["model_name"],
"model_info": model_info,
}


@csrf_exempt
@require_GET
@token_authenticated(token_auth_only=True)
Expand Down Expand Up @@ -42,3 +63,32 @@ async def models(
"object": "list",
}
)


@csrf_exempt
@require_GET
@token_authenticated(token_auth_only=True)
@tos_accepted
@log_request
async def model_group_info(
request: ASGIRequest, token: Token, request_log: Request, *args: Any, **kwargs: Any
) -> JsonResponse:
"""LiteLLM-style rich model metadata endpoint.

Returns one entry per configured model (JSON array) so that LiteLLM-aware
clients — which probe /model_group/info before falling back to
/v1/models — read token limits and capabilities directly from the gateway
instead of guessing from bundled model catalogs.
"""
router_config = get_router_config()
model_list: list[dict[str, Any]] = router_config["model_list"]
excluded_models = set(await sync_to_async(token.model_exclusion_list)())

return JsonResponse(
[
_model_group_info_entry(model)
for model in model_list
if model["model_name"] not in excluded_models
],
safe=False,
)
Loading