Skip to content
Merged
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
20 changes: 16 additions & 4 deletions src/any_llm/providers/fireworks/fireworks.py
Original file line number Diff line number Diff line change
@@ -1,4 +1,4 @@
from collections.abc import AsyncGenerator, AsyncIterator, Iterator
from collections.abc import AsyncGenerator, AsyncIterator, Iterator, Sequence
from typing import Any, cast

try:
Expand All @@ -22,20 +22,22 @@
CompletionUsage,
Reasoning,
)
from any_llm.types.model import Model
from any_llm.types.responses import Response, ResponseStreamEvent


class FireworksProvider(Provider):
PROVIDER_NAME = "fireworks"
ENV_API_KEY_NAME = "FIREWORKS_API_KEY"
PROVIDER_DOCUMENTATION_URL = "https://fireworks.ai/api"
BASE_URL = "https://api.fireworks.ai/inference/v1"

SUPPORTS_COMPLETION_STREAMING = True
SUPPORTS_COMPLETION = True
SUPPORTS_RESPONSES = True
SUPPORTS_COMPLETION_REASONING = False
SUPPORTS_EMBEDDING = False
SUPPORTS_LIST_MODELS = False
SUPPORTS_LIST_MODELS = True

PACKAGES_INSTALLED = PACKAGES_INSTALLED

Expand Down Expand Up @@ -194,7 +196,7 @@ async def aresponses(
) -> Response | AsyncIterator[ResponseStreamEvent]:
"""Call Fireworks Responses API and normalize into ChatCompletion/Chunks."""
client = AsyncOpenAI(
base_url="https://api.fireworks.ai/inference/v1",
base_url=self.BASE_URL,
api_key=self.config.api_key,
)
response = await client.responses.create(
Expand All @@ -218,7 +220,7 @@ async def aresponses(
def responses(self, model: str, input_data: Any, **kwargs: Any) -> Response | Iterator[ResponseStreamEvent]:
"""Call Fireworks Responses API and normalize into ChatCompletion/Chunks."""
client = OpenAI(
base_url="https://api.fireworks.ai/inference/v1",
base_url=self.BASE_URL,
api_key=self.config.api_key,
)
response = client.responses.create(
Expand All @@ -238,3 +240,13 @@ def responses(self, model: str, input_data: Any, **kwargs: Any) -> Response | It
response.reasoning = Reasoning(content=reasoning) if reasoning else None # type: ignore[assignment]

return response

def list_models(self, **kwargs: Any) -> Sequence[Model]:
"""
Fetch available models from the /v1/models endpoint.
"""
client = OpenAI(
base_url=self.BASE_URL,
api_key=self.config.api_key,
)
return client.models.list(**kwargs).data