diff --git a/py/packages/genkit-openai/src/genkit_openai/models/model.py b/py/packages/genkit-openai/src/genkit_openai/models/model.py index 55d9cf0219..ed474be71d 100644 --- a/py/packages/genkit-openai/src/genkit_openai/models/model.py +++ b/py/packages/genkit-openai/src/genkit_openai/models/model.py @@ -56,7 +56,22 @@ _GENKIT_ONLY = frozenset({'api_key', 'top_k', 'version', 'max_output_tokens', 'stop_sequences'}) -def _openai_create_kwargs(*, config: OpenAIConfig) -> dict[str, Any]: +def _uses_max_completion_tokens(model: str | None) -> bool: + """Whether a model requires the reasoning-model token-limit parameter.""" + if not model: + return False + model_id = model.rsplit('/', 1)[-1].lower() + # Fine-tuned OpenAI model ids are prefixed with ``ft:`` and custom + # deployment names may put the base model after an arbitrary prefix. + if model_id.startswith('ft:'): + parts = model_id.split(':') + if len(parts) > 1: + model_id = parts[1] + reasoning_prefixes = ('o1', 'o3', 'o4', 'gpt-5', 'gpt-6') + return model_id.startswith(reasoning_prefixes) or any(f'-{prefix}' in model_id for prefix in reasoning_prefixes) + + +def _openai_create_kwargs(*, config: OpenAIConfig, model: str | None = None) -> dict[str, Any]: """Kwargs for chat.completions.create(). Peel Genkit-only keys. ``stop_sequences`` becomes ``stop`` when ``stop`` @@ -64,6 +79,8 @@ def _openai_create_kwargs(*, config: OpenAIConfig) -> dict[str, Any]: caller when ``OpenAIConfig.model`` is unset. Everything else, including extras, goes out under the Python field name. ``max_output_tokens`` is not mapped to ``max_tokens`` — that knob is ``max_tokens`` / ``maxTokens``. + For reasoning models, ``max_tokens`` is emitted as ``max_completion_tokens`` + because the OpenAI API rejects the deprecated field. """ body: dict[str, Any] = {} for name in type(config).model_fields: @@ -71,6 +88,15 @@ def _openai_create_kwargs(*, config: OpenAIConfig) -> dict[str, Any]: continue value = getattr(config, name) if value is not None: + if name == 'max_tokens': + # OpenAI reasoning models reject the deprecated max_tokens + # field. Keep the explicit max_completion_tokens value when + # both knobs are supplied so the request remains valid. + if config.max_completion_tokens is not None: + continue + if _uses_max_completion_tokens(model) or config.reasoning_effort is not None: + body['max_completion_tokens'] = value + continue body[name] = value extras = config.model_extra if extras: @@ -382,7 +408,8 @@ async def _get_openai_request_config(self, request: ModelRequest) -> dict: ) if config.version: openai_config['model'] = config.version - openai_config.update(_openai_create_kwargs(config=config)) + effective_model = config.model or config.version or self._model + openai_config.update(_openai_create_kwargs(config=config, model=effective_model)) return openai_config async def _generate(self, request: ModelRequest) -> ModelResponse: diff --git a/py/packages/genkit-openai/tests/openai_model_test.py b/py/packages/genkit-openai/tests/openai_model_test.py index cc27c8c825..7ec2f4fc42 100644 --- a/py/packages/genkit-openai/tests/openai_model_test.py +++ b/py/packages/genkit-openai/tests/openai_model_test.py @@ -25,7 +25,7 @@ from genkit_openai.models import OpenAIModel from genkit_openai.models.model import _usage_from_completion from genkit_openai.models.utils import strip_markdown_fences -from genkit_openai.typing import OpenAIConfig +from genkit_openai.typing import OpenAIConfig, ReasoningEffort from openai.types import CompletionUsage from openai.types.chat import ChatCompletionChunk from pydantic import BaseModel @@ -131,6 +131,62 @@ async def test_get_openai_config_peels_genkit_keys_and_passes_the_rest() -> None assert 'version' not in body +@pytest.mark.asyncio +@pytest.mark.parametrize( + ('model_name', 'reasoning_effort'), + [ + ('gpt-6-astra', None), + ('ft:o1-mini:my-org:custom', None), + ('my-o1-mini-deployment', None), + ('my-deployment', ReasoningEffort.HIGH), + ], +) +async def test_get_openai_config_uses_max_completion_tokens_for_reasoning_models( + model_name: str, reasoning_effort: ReasoningEffort | None +) -> None: + """Reasoning models reject the deprecated max_tokens request field.""" + model = OpenAIModel(model=model_name, client=MagicMock()) + request = ModelRequest( + messages=[Message(role=Role.USER, content=[Part(root=TextPart(text='hi'))])], + config=OpenAIConfig(max_tokens=32, reasoning_effort=reasoning_effort), + ) + + body = await model._get_openai_request_config(request) + + assert body['max_completion_tokens'] == 32 + assert 'max_tokens' not in body + + +@pytest.mark.asyncio +async def test_get_openai_config_keeps_max_tokens_for_legacy_models() -> None: + """Legacy OpenAI-compatible models continue to receive max_tokens.""" + model = OpenAIModel(model='gpt-4o', client=MagicMock()) + request = ModelRequest( + messages=[Message(role=Role.USER, content=[Part(root=TextPart(text='hi'))])], + config=OpenAIConfig(max_tokens=32), + ) + + body = await model._get_openai_request_config(request) + + assert body['max_tokens'] == 32 + assert 'max_completion_tokens' not in body + + +@pytest.mark.asyncio +async def test_get_openai_config_prefers_explicit_max_completion_tokens() -> None: + """An explicit modern token limit wins when both fields are configured.""" + model = OpenAIModel(model='gpt-4o', client=MagicMock()) + request = ModelRequest( + messages=[Message(role=Role.USER, content=[Part(root=TextPart(text='hi'))])], + config=OpenAIConfig(max_tokens=32, max_completion_tokens=64), + ) + + body = await model._get_openai_request_config(request) + + assert body['max_completion_tokens'] == 64 + assert 'max_tokens' not in body + + @pytest.mark.asyncio async def test_get_openai_config_model_field_overrides_version() -> None: """OpenAIConfig.model is the create() model id; it wins over version."""