diff --git a/openjudge/models/qwen_vl_model.py b/openjudge/models/qwen_vl_model.py index 06ef2c6dd..a240df64e 100644 --- a/openjudge/models/qwen_vl_model.py +++ b/openjudge/models/qwen_vl_model.py @@ -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 @@ -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) @@ -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 @@ -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 diff --git a/tests/models/test_qwen_vl_model.py b/tests/models/test_qwen_vl_model.py new file mode 100644 index 000000000..758bde28b --- /dev/null +++ b/tests/models/test_qwen_vl_model.py @@ -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