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: 19 additions & 8 deletions openjudge/models/qwen_vl_model.py
Original file line number Diff line number Diff line change
Expand Up @@ -49,6 +49,7 @@ def __init__(
temperature: float = 0.1,
top_p: float = 0.9,
max_tokens: int = 2000,
timeout: Optional[float] = None,
):
"""
Initialize Qwen VL API client
Expand All @@ -59,6 +60,8 @@ def __init__(
temperature: Sampling temperature
top_p: Nucleus sampling
max_tokens: Maximum tokens to generate
timeout: Request timeout in seconds (defaults to the DashScope
SDK's own default of 300s when not set)
"""
super().__init__(model=model, stream=False)

Expand All @@ -72,6 +75,7 @@ def __init__(
self.temperature = temperature
self.top_p = top_p
self.max_tokens = max_tokens
self.timeout = timeout

# Cost tracking
self._total_requests = 0
Expand Down Expand Up @@ -151,15 +155,22 @@ def generate(
messages = self._format_messages(content, system_prompt)

# Call API
call_kwargs: Dict[str, Any] = {
"api_key": self.api_key,
"model": self.model,
"messages": messages,
"temperature": self.temperature,
"top_p": self.top_p,
"max_length": self.max_tokens,
}
if self.timeout is not None:
# DashScope reads the socket timeout from `request_timeout`;
# a `timeout=` kwarg is accepted but dropped into the request
# body unused, so this is not a naming choice.
call_kwargs["request_timeout"] = self.timeout

try:
response = MultiModalConversation.call(
api_key=self.api_key,
model=self.model,
messages=messages,
temperature=self.temperature,
top_p=self.top_p,
max_length=self.max_tokens,
)
response = MultiModalConversation.call(**call_kwargs)

self._total_requests += 1

Expand Down
52 changes: 52 additions & 0 deletions tests/models/test_qwen_vl_model.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,52 @@
# -*- coding: utf-8 -*-
"""Unit tests for QwenVLModel."""

from unittest.mock import MagicMock, patch

import pytest

from openjudge.models.qwen_vl_model import QwenVLModel


def _fake_response(text: str = "OK"):
response = MagicMock()
response.status_code = 200
response.message = ""
response.output.choices = [MagicMock()]
response.output.choices[0].message.content = [{"text": text}]
return response


@pytest.mark.unit
class TestQwenVLModelTimeout:
"""Test cases for QwenVLModel's timeout handling."""

@pytest.mark.parametrize(
"init_kwargs, expect_request_timeout",
[
({"timeout": 5.0}, 5.0),
({}, None),
({"timeout": None}, None),
],
ids=["with_timeout", "defaults", "explicit_none"],
)
@patch("openjudge.models.qwen_vl_model.MultiModalConversation")
def test_generate_forwards_request_timeout(
self,
mock_conversation,
init_kwargs,
expect_request_timeout,
):
"""timeout=N must reach DashScope as request_timeout=N, never as timeout=N."""
mock_conversation.call.return_value = _fake_response()

model = QwenVLModel(api_key="test-key", **init_kwargs)
model.generate(text="hi")

call_kwargs = mock_conversation.call.call_args[1]
assert "timeout" not in call_kwargs

if expect_request_timeout is None:
assert "request_timeout" not in call_kwargs
else:
assert call_kwargs["request_timeout"] == expect_request_timeout