diff --git a/docs/ai-python/content/docs/basics/streaming.mdx b/docs/ai-python/content/docs/basics/streaming.mdx index 6a3a7c77..d579cdce 100644 --- a/docs/ai-python/content/docs/basics/streaming.mdx +++ b/docs/ai-python/content/docs/basics/streaming.mdx @@ -136,15 +136,3 @@ async with ai.stream(model, messages) as stream: for f in stream.message.files: print(f.media_type) ``` - -Use `ai.generate` for dedicated image and video models: - -```python -result = await ai.generate( - model, - [ai.user_message("A watercolor mothership over a quiet city.")], - ai.ImageParams(n=1, aspect_ratio="16:9"), -) - -image = result.images[0] -``` diff --git a/docs/ai-python/content/docs/reference/ai/generate.mdx b/docs/ai-python/content/docs/reference/ai/generate.mdx deleted file mode 100644 index 79d780d8..00000000 --- a/docs/ai-python/content/docs/reference/ai/generate.mdx +++ /dev/null @@ -1,35 +0,0 @@ ---- -title: "generate" -description: Generate non-streaming media responses. -type: reference -summary: Reference for ai.generate. ---- - -`ai.generate` calls non-streaming generation APIs, such as image and -video models, and returns a `Message`. - -```python -result = await ai.generate( - model, - [ai.user_message("A watercolor mothership over a quiet city.")], - ai.ImageParams(n=1, aspect_ratio="16:9"), -) -``` - -## Function - -```python -await ai.generate(model, messages, params) -``` - -## Arguments - -- `model`: `ai.Model`. -- `messages`: list of `ai.messages.Message`. -- `params`: `ai.ImageParams` or `ai.VideoParams`. - -## Return value - -`ai.generate` returns an assistant `Message`. Read generated files through -`message.files` or use the normal message helpers for text, reasoning, and -metadata. diff --git a/docs/ai-python/content/docs/reference/ai/index.mdx b/docs/ai-python/content/docs/reference/ai/index.mdx index 19402a9b..a6f3fe3c 100644 --- a/docs/ai-python/content/docs/reference/ai/index.mdx +++ b/docs/ai-python/content/docs/reference/ai/index.mdx @@ -21,8 +21,6 @@ These top-level exports have their own pages under `ai`. - [`stream`](/docs/reference/ai/stream): Call a model and iterate response events. The same page documents `Stream`. -- [`generate`](/docs/reference/ai/generate): Call non-streaming media - generation models. - [`get_model`](/docs/reference/ai/get-model): Resolve a model id into a `Model`. - [`get_provider`](/docs/reference/ai/get-provider): Resolve and configure a @@ -124,8 +122,7 @@ await protocol.generate(client, model, messages, params, provider=provider.name) Model params are top-level `ai` types. -Use `InferenceRequestParams` with `stream` and `Agent.run`. Use `ImageParams` -and `VideoParams` with `generate`. +Use `InferenceRequestParams` with `stream` and `Agent.run`. ```python params = ai.InferenceRequestParams().with_temperature(0) @@ -176,28 +173,6 @@ Routing params: - `ProviderRankingStrategy` - `GLOBAL` -Media generation params: - -```python -ai.ImageParams( - n=1, - size=None, - aspect_ratio=None, - seed=None, - provider_options={}, -) - -ai.VideoParams( - n=1, - aspect_ratio=None, - resolution=None, - duration=None, - fps=None, - seed=None, - provider_options={}, -) -``` - ## Messages Message builders create `Message` values. diff --git a/docs/ai-python/content/docs/reference/ai/meta.json b/docs/ai-python/content/docs/reference/ai/meta.json index 56627c36..f878e7bf 100644 --- a/docs/ai-python/content/docs/reference/ai/meta.json +++ b/docs/ai-python/content/docs/reference/ai/meta.json @@ -3,7 +3,6 @@ "description": "Reference for top-level ai exports.", "pages": [ "stream", - "generate", "get-model", "get-provider", "tool-decorator", diff --git a/examples/.test_scripts/run-examples.py b/examples/.test_scripts/run-examples.py index b5afb081..a7d35890 100755 --- a/examples/.test_scripts/run-examples.py +++ b/examples/.test_scripts/run-examples.py @@ -13,7 +13,7 @@ uv run examples/.test_scripts/run-examples.py --model MODEL # patch ai.get_model() to use the given model for every sample uv run examples/.test_scripts/run-examples.py --protocol=responses - # patch model/provider helpers and ai.stream()/ai.generate() + # patch model/provider helpers, ai.stream(), and experimental_generate() """ import argparse diff --git a/examples/.test_scripts/run-with-patched-model.py b/examples/.test_scripts/run-with-patched-model.py index b716495d..33c0b4a8 100644 --- a/examples/.test_scripts/run-with-patched-model.py +++ b/examples/.test_scripts/run-with-patched-model.py @@ -86,7 +86,7 @@ def main() -> None: original_get_model = _model.get_model original_stream = _api.stream - original_generate = _api.generate + original_generate = _api.experimental_generate def selected_protocol() -> ai.ProviderProtocol[Any] | None: if protocol_factory is None: @@ -203,10 +203,7 @@ def __init__( cast("Any", core).stream = patched_stream cast("Any", _api).stream = patched_stream - cast("Any", ai).generate = patched_generate - cast("Any", models).generate = patched_generate - cast("Any", core).generate = patched_generate - cast("Any", _api).generate = patched_generate + cast("Any", _api).experimental_generate = patched_generate sys.argv = [args.file] runpy.run_path(args.file, run_name="__main__") diff --git a/examples/media/image_edit.py b/examples/media/image_edit.py index 1f3b2ae4..4b5b7f45 100644 --- a/examples/media/image_edit.py +++ b/examples/media/image_edit.py @@ -10,6 +10,7 @@ import pathlib import ai +from ai.models.core import api, params model = ai.get_model("openai/gpt-image-1") @@ -32,8 +33,8 @@ async def main() -> None: ), ] - result = await ai.generate( - model, messages, ai.ImageParams(size="1024x1024") + result = await api.experimental_generate( + model, messages, params.ImageParams(size="1024x1024") ) print(f"Generated {len(result.images)} edited image(s)") diff --git a/examples/media/image_generation.py b/examples/media/image_generation.py index 4817630f..8ced3892 100644 --- a/examples/media/image_generation.py +++ b/examples/media/image_generation.py @@ -1,10 +1,11 @@ -"""Image generation — dedicated image model via generate().""" +"""Image generation — dedicated image model via experimental_generate().""" import asyncio import base64 import pathlib import ai +from ai.models.core import api, params model = ai.get_model("google/imagen-4.0-generate-001") @@ -18,8 +19,8 @@ async def main() -> None: - result = await ai.generate( - model, messages, ai.ImageParams(n=2, aspect_ratio="16:9") + result = await api.experimental_generate( + model, messages, params.ImageParams(n=2, aspect_ratio="16:9") ) print(f"Generated {len(result.images)} image(s)") diff --git a/examples/media/video_generation.py b/examples/media/video_generation.py index 2605d70d..1d4befcb 100644 --- a/examples/media/video_generation.py +++ b/examples/media/video_generation.py @@ -1,10 +1,11 @@ -"""Video generation — dedicated video model via generate().""" +"""Video generation — dedicated video model via experimental_generate().""" import asyncio import base64 import pathlib import ai +from ai.models.core import api, params model = ai.get_model("google/veo-3.0-generate-001") @@ -19,10 +20,10 @@ async def main() -> None: print("Generating video (this may take a minute or two)...") - result = await ai.generate( + result = await api.experimental_generate( model, messages, - ai.VideoParams(aspect_ratio="16:9", duration=8), + params.VideoParams(aspect_ratio="16:9", duration=8), ) print(f"Generated {len(result.videos)} video(s)") diff --git a/src/ai/__init__.py b/src/ai/__init__.py index ab9b3e58..9206b427 100644 --- a/src/ai/__init__.py +++ b/src/ai/__init__.py @@ -58,7 +58,6 @@ CloudRegion, ContextManagementParams, GeoRegion, - ImageParams, InferenceRequestParams, MinPSamplerParams, Model, @@ -85,8 +84,6 @@ TopKSamplerParams, TopPSamplerParams, Unset, - VideoParams, - generate, get_model, probe, stream, @@ -123,7 +120,6 @@ "HTTPErrorContext", "HookDeferredException", "HookRegistry", - "ImageParams", "InferenceRequestParams", "InstallationError", "MinPSamplerParams", @@ -178,7 +174,6 @@ "TopPSamplerParams", "Unset", "UnsupportedProviderError", - "VideoParams", "assistant_message", "cancel_hook", "content_output", @@ -187,7 +182,6 @@ "errors", "events", "file_part", - "generate", "get_hook_registry", "get_model", "get_provider", diff --git a/src/ai/models/__init__.py b/src/ai/models/__init__.py index a07f85ae..2daf6379 100644 --- a/src/ai/models/__init__.py +++ b/src/ai/models/__init__.py @@ -35,7 +35,6 @@ from ..providers.base import Provider, ProviderProtocol from .core.api import ( Stream, - generate, probe, stream, ) @@ -48,9 +47,7 @@ CacheParams, CloudRegion, ContextManagementParams, - GenerateParams, GeoRegion, - ImageParams, InferenceRequestParams, MinPSamplerParams, ModelProviderDefault, @@ -73,7 +70,6 @@ TopKSamplerParams, TopPSamplerParams, Unset, - VideoParams, ) __all__ = [ @@ -84,9 +80,7 @@ "CacheParams", "CloudRegion", "ContextManagementParams", - "GenerateParams", "GeoRegion", - "ImageParams", "InferenceRequestParams", "MinPSamplerParams", "Model", @@ -113,8 +107,6 @@ "TopKSamplerParams", "TopPSamplerParams", "Unset", - "VideoParams", - "generate", "get_model", "probe", "stream", diff --git a/src/ai/models/core/__init__.py b/src/ai/models/core/__init__.py index a4e01e20..51cb1b49 100644 --- a/src/ai/models/core/__init__.py +++ b/src/ai/models/core/__init__.py @@ -4,7 +4,6 @@ from . import helpers from .api import ( Stream, - generate, probe, stream, ) @@ -17,9 +16,7 @@ CacheParams, CloudRegion, ContextManagementParams, - GenerateParams, GeoRegion, - ImageParams, InferenceRequestParams, MinPSamplerParams, ModelProviderDefault, @@ -42,7 +39,6 @@ TopKSamplerParams, TopPSamplerParams, Unset, - VideoParams, ) __all__ = [ @@ -53,9 +49,7 @@ "CacheParams", "CloudRegion", "ContextManagementParams", - "GenerateParams", "GeoRegion", - "ImageParams", "InferenceRequestParams", "MinPSamplerParams", "Model", @@ -81,8 +75,6 @@ "TopKSamplerParams", "TopPSamplerParams", "Unset", - "VideoParams", - "generate", "get_model", "helpers", "probe", diff --git a/src/ai/models/core/api.py b/src/ai/models/core/api.py index fdbcd942..6faeccfb 100644 --- a/src/ai/models/core/api.py +++ b/src/ai/models/core/api.py @@ -606,12 +606,15 @@ async def _stream( await s.aclose() -async def generate( +async def experimental_generate( model: model_.Model, messages: list[types.messages.Message], params: params_.GenerateParams, ) -> types.messages.Message: - """Generate a non-streaming response (images, video, etc.).""" + """Generate a non-streaming response (images, video, etc.). + + Experimental: not part of the stable API, may change or be removed. + """ request = _GenerateRequest(model, list(messages), params) async with telemetry.span( telemetry.AiGenerateSpanData( diff --git a/src/ai/providers/ai_gateway/protocol.py b/src/ai/providers/ai_gateway/protocol.py index 69e6deda..5b505ec7 100644 --- a/src/ai/providers/ai_gateway/protocol.py +++ b/src/ai/providers/ai_gateway/protocol.py @@ -1060,7 +1060,7 @@ async def _generate_image( gateway: gateway_client.GatewayClient, model: core.model.Model, messages: list[types.messages.Message], - params: core.ImageParams, + params: params_.ImageParams, ) -> types.messages.Message: """Hit ``/image-model`` and return a Message with FileParts.""" prompt = _extract_prompt(messages) @@ -1101,7 +1101,7 @@ async def _generate_video( gateway: gateway_client.GatewayClient, model: core.model.Model, messages: list[types.messages.Message], - params: core.VideoParams, + params: params_.VideoParams, ) -> types.messages.Message: """Hit ``/video-model`` (SSE) and return a Message with FileParts.""" prompt = _extract_prompt(messages) @@ -1168,11 +1168,11 @@ async def generate( gateway: gateway_client.GatewayClient, model: core.model.Model, messages: list[types.messages.Message], - params: core.GenerateParams, + params: params_.GenerateParams, ) -> types.messages.Message: """Generate media through the AI Gateway.""" try: - if isinstance(params, core.VideoParams): + if isinstance(params, params_.VideoParams): return await _generate_video(gateway, model, messages, params) return await _generate_image(gateway, model, messages, params) except client_errors.GatewayError as exc: @@ -1210,7 +1210,7 @@ async def generate( client: gateway_client.GatewayClient, model: core.model.Model, messages: list[types.messages.Message], - params: core.GenerateParams, + params: params_.GenerateParams, *, provider: str, ) -> types.messages.Message: diff --git a/tests/conftest.py b/tests/conftest.py index 1e776873..4f34d11a 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -8,6 +8,7 @@ import ai from ai import models +from ai.models.core import params from ai.providers import history_utils from ai.types import builders from ai.types import events as agent_events_ @@ -71,7 +72,7 @@ async def generate( self, model: models.Model, messages: list[messages_.Message], - params: models.GenerateParams, + params: params.GenerateParams, ) -> messages_.Message: if model.protocol is not None: return await model.protocol.generate( @@ -229,7 +230,7 @@ async def generate( self, model: models.Model, messages: list[messages_.Message], - params: models.GenerateParams, + params: params.GenerateParams, ) -> messages_.Message: if self._call_index >= len(self._responses): raise RuntimeError( diff --git a/tests/models/core/test_api.py b/tests/models/core/test_api.py index 604cd214..d9365120 100644 --- a/tests/models/core/test_api.py +++ b/tests/models/core/test_api.py @@ -8,6 +8,7 @@ import ai from ai import models +from ai.models.core import api, params from ai.types import events as events_ from ai.types import messages as messages_ @@ -338,7 +339,7 @@ async def test_generate_dispatches_to_provider() -> None: async def _generate( model: models.Model, messages: list[messages_.Message], - params: models.GenerateParams, + params: params.GenerateParams, ) -> messages_.Message: nonlocal called called = True @@ -346,10 +347,10 @@ async def _generate( provider._generate_impl = _generate - result = await models.generate( + result = await api.experimental_generate( model, [ai.user_message("A cat")], - models.ImageParams(n=1), + params.ImageParams(n=1), ) assert called @@ -372,17 +373,17 @@ async def generate( client: Any, model: models.Model, messages: list[messages_.Message], - params: models.GenerateParams, + params: params.GenerateParams, *, provider: str, ) -> messages_.Message: _ = client, model, messages, params, provider return sentinel - result = await models.generate( + result = await api.experimental_generate( MOCK_MODEL.with_protocol(OverrideProtocol()), [ai.user_message("A cat")], - models.ImageParams(n=1), + params.ImageParams(n=1), ) assert result is sentinel diff --git a/tests/providers/ai_gateway/test_generate_image.py b/tests/providers/ai_gateway/test_generate_image.py index 2f751e93..86ec8470 100644 --- a/tests/providers/ai_gateway/test_generate_image.py +++ b/tests/providers/ai_gateway/test_generate_image.py @@ -22,6 +22,7 @@ import pytest import ai +from ai.models.core import api from ai.models.core.params import ImageParams from ai.types import messages @@ -56,7 +57,7 @@ def handler(req: httpx.Request) -> httpx.Response: model = mock_model( httpx.MockTransport(handler), model_id=_IMAGE_MODEL_ID ) - msg = await ai.generate( + msg = await api.experimental_generate( model, [user_msg("A sunset over Tokyo")], ImageParams() ) @@ -76,7 +77,7 @@ def handler(req: httpx.Request) -> httpx.Response: json={"images": [_PNG_B64, _JPEG_B64, _PNG_B64]}, ) - msg = await ai.generate( + msg = await api.experimental_generate( mock_model(httpx.MockTransport(handler), model_id=_IMAGE_MODEL_ID), [user_msg("Three cats")], params=ImageParams(n=3), @@ -99,7 +100,7 @@ def handler(req: httpx.Request) -> httpx.Response: }, ) - msg = await ai.generate( + msg = await api.experimental_generate( mock_model(httpx.MockTransport(handler), model_id=_IMAGE_MODEL_ID), [user_msg("a dog")], ImageParams(), @@ -128,7 +129,7 @@ def handler(req: httpx.Request) -> httpx.Response: api_key="sk-test", model_id="openai/gpt-image-1", ) - await ai.generate(model, [user_msg("Hi")], ImageParams()) + await api.experimental_generate(model, [user_msg("Hi")], ImageParams()) assert captured["authorization"] == "Bearer sk-test" assert captured["ai-image-model-specification-version"] == "3" @@ -149,7 +150,7 @@ def handler(req: httpx.Request) -> httpx.Response: messages.FilePart(data=_PNG_B64, media_type="image/png"), ], ) - await ai.generate( + await api.experimental_generate( mock_model(httpx.MockTransport(handler), model_id=_IMAGE_MODEL_ID), [msg], params=ImageParams( @@ -181,7 +182,7 @@ def handler(req: httpx.Request) -> httpx.Response: captured_url.append(str(req.url)) return httpx.Response(200, json={"images": [_PNG_B64]}) - await ai.generate( + await api.experimental_generate( mock_model(httpx.MockTransport(handler), model_id=_IMAGE_MODEL_ID), [user_msg("test")], ImageParams(), @@ -209,7 +210,7 @@ def handler(req: httpx.Request) -> httpx.Response: ) with pytest.raises(ai.ProviderAuthenticationError): - await ai.generate( + await api.experimental_generate( mock_model( httpx.MockTransport(handler), model_id=_IMAGE_MODEL_ID ), @@ -230,7 +231,7 @@ def handler(req: httpx.Request) -> httpx.Response: ) with pytest.raises(ai.ProviderRateLimitError): - await ai.generate( + await api.experimental_generate( mock_model( httpx.MockTransport(handler), model_id=_IMAGE_MODEL_ID ), @@ -244,7 +245,7 @@ async def test_empty_images_returns_empty_message(self) -> None: def handler(req: httpx.Request) -> httpx.Response: return httpx.Response(200, json={"images": []}) - msg = await ai.generate( + msg = await api.experimental_generate( mock_model(httpx.MockTransport(handler), model_id=_IMAGE_MODEL_ID), [user_msg("test")], ImageParams(), diff --git a/tests/providers/ai_gateway/test_generate_video.py b/tests/providers/ai_gateway/test_generate_video.py index 04d2a7a3..ee96ae88 100644 --- a/tests/providers/ai_gateway/test_generate_video.py +++ b/tests/providers/ai_gateway/test_generate_video.py @@ -23,6 +23,7 @@ import pytest import ai +from ai.models.core import api from ai.models.core.params import VideoParams from ai.types import messages @@ -65,7 +66,7 @@ async def test_basic_video_generation_base64(self) -> None: def handler(req: httpx.Request) -> httpx.Response: return httpx.Response(200, text=body) - msg = await ai.generate( + msg = await api.experimental_generate( mock_model(httpx.MockTransport(handler), model_id=_VIDEO_MODEL_ID), [user_msg("A cat walking on a beach")], params=VideoParams(), @@ -103,7 +104,7 @@ def handler(req: httpx.Request) -> httpx.Response: new_callable=AsyncMock, return_value=(_MP4_HEADER, "video/mp4"), ) as mock_dl: - msg = await ai.generate( + msg = await api.experimental_generate( model, [user_msg("A sunset timelapse")], params=VideoParams(), @@ -136,7 +137,7 @@ async def test_multiple_videos(self) -> None: def handler(req: httpx.Request) -> httpx.Response: return httpx.Response(200, text=body) - msg = await ai.generate( + msg = await api.experimental_generate( mock_model(httpx.MockTransport(handler), model_id=_VIDEO_MODEL_ID), [user_msg("Two versions")], params=VideoParams(n=2), @@ -178,7 +179,7 @@ def handler(req: httpx.Request) -> httpx.Response: api_key="sk-test", model_id=_VIDEO_MODEL_ID, ) - await ai.generate( + await api.experimental_generate( model, [user_msg("test")], params=VideoParams(), @@ -221,7 +222,7 @@ def handler(req: httpx.Request) -> httpx.Response: messages.FilePart(data=png_b64, media_type="image/png"), ], ) - await ai.generate( + await api.experimental_generate( mock_model(httpx.MockTransport(handler), model_id=_VIDEO_MODEL_ID), [msg], params=VideoParams( @@ -270,7 +271,7 @@ def handler(req: httpx.Request) -> httpx.Response: ), ) - await ai.generate( + await api.experimental_generate( mock_model(httpx.MockTransport(handler), model_id=_VIDEO_MODEL_ID), [user_msg("test")], params=VideoParams(), @@ -300,7 +301,7 @@ def handler(req: httpx.Request) -> httpx.Response: return httpx.Response(200, text=body) with pytest.raises(ai.ProviderBadRequestError, match="Content policy"): - await ai.generate( + await api.experimental_generate( mock_model( httpx.MockTransport(handler), model_id=_VIDEO_MODEL_ID ), @@ -321,7 +322,7 @@ def handler(req: httpx.Request) -> httpx.Response: ) with pytest.raises(ai.ProviderAuthenticationError): - await ai.generate( + await api.experimental_generate( mock_model( httpx.MockTransport(handler), model_id=_VIDEO_MODEL_ID ), @@ -336,7 +337,7 @@ def handler(req: httpx.Request) -> httpx.Response: return httpx.Response(200, text="") with pytest.raises(ai.ProviderResponseError, match="SSE stream ended"): - await ai.generate( + await api.experimental_generate( mock_model( httpx.MockTransport(handler), model_id=_VIDEO_MODEL_ID ),