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
31 changes: 29 additions & 2 deletions py/packages/genkit-openai/src/genkit_openai/models/model.py
Original file line number Diff line number Diff line change
Expand Up @@ -56,21 +56,47 @@
_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``
was not set. ``version`` is peeled here and applied as ``model`` by the
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:
if name in _GENKIT_ONLY:
continue
value = getattr(config, name)
if value is not None:
if name == 'max_tokens':
Comment thread
hilariie marked this conversation as resolved.
# 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:
Comment thread
hilariie marked this conversation as resolved.
body['max_completion_tokens'] = value
continue
body[name] = value
extras = config.model_extra
if extras:
Expand Down Expand Up @@ -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:
Expand Down
58 changes: 57 additions & 1 deletion py/packages/genkit-openai/tests/openai_model_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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

Comment thread
hilariie marked this conversation as resolved.

@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."""
Expand Down
Loading