feat(fireworks): add document reranking (#39732)

This commit is contained in:
Noah Dylan authored and GitHub committed 2026-08-20 12:57:26 -04:00
1 parent 8df1265122
commit b8d1ab5946
4 files changed
+459

No files matched your search

@@ -4,10 +4,12 @@ from langchain_fireworks._version import __version__
from langchain_fireworks.chat_models import ChatFireworks
from langchain_fireworks.embeddings import FireworksEmbeddings
from langchain_fireworks.llms import Fireworks
from langchain_fireworks.rerank import FireworksRerank
__all__ = [
"ChatFireworks",
"Fireworks",
"FireworksEmbeddings",
"FireworksRerank",
"__version__",
]
@@ -0,0 +1,239 @@
"""Fireworks document reranking integration."""
from __future__ import annotations
import json
from collections.abc import Mapping, Sequence
from copy import deepcopy
from typing import Any
from langchain_core._api import beta
from langchain_core.callbacks import Callbacks
from langchain_core.documents import BaseDocumentCompressor, Document
from langchain_core.utils import secret_from_env
from openai import AsyncOpenAI, OpenAI
from pydantic import ConfigDict, Field, SecretStr, model_validator
from typing_extensions import override
# The OpenAI SDK unpacks `get_args()` on a `dict` `cast_to`, so a bare `dict`
# raises `ValueError` while parsing the response. Keep this parameterized.
_RESPONSE_TYPE = dict[str, Any]
@beta()
class FireworksRerank(BaseDocumentCompressor):
"""Document compressor that uses Fireworks' reranking API."""
client: Any = None
"""OpenAI-compatible client used to call Fireworks."""
async_client: Any = None
"""Async OpenAI-compatible client used to call Fireworks."""
top_n: int | None = 3
"""Number of documents to return."""
model: str
"""Fireworks reranking model to use."""
fireworks_api_key: SecretStr | None = Field(
default_factory=secret_from_env("FIREWORKS_API_KEY", default=None)
)
"""Fireworks API key."""
base_url: str = "https://api.fireworks.ai/inference/v1"
"""Base URL for the Fireworks API."""
user_agent: str = "langchain:partner"
"""Identifier for the application making the request."""
model_config = ConfigDict(
arbitrary_types_allowed=True,
extra="forbid",
)
@model_validator(mode="after")
def validate_environment(self) -> FireworksRerank:
"""Create the OpenAI-compatible clients that were not supplied."""
if self.client is None or self.async_client is None:
if self.fireworks_api_key is None:
msg = (
"FIREWORKS_API_KEY is required unless both client and "
"async_client are supplied."
)
raise ValueError(msg)
client_kwargs: dict[str, Any] = {
"api_key": self.fireworks_api_key.get_secret_value(),
"base_url": self.base_url,
"default_headers": {"User-Agent": self.user_agent},
}
if self.client is None:
self.client = OpenAI(**client_kwargs)
if self.async_client is None:
self.async_client = AsyncOpenAI(**client_kwargs)
return self
def _document_to_str(
self,
document: str | Document | Mapping[str, Any],
rank_fields: Sequence[str] | None = None,
) -> str:
"""Convert a supported document value to the string API format."""
if isinstance(document, Document):
return document.page_content
if isinstance(document, Mapping):
value: Mapping[str, Any] = document
if rank_fields is not None:
value = {key: document[key] for key in rank_fields if key in document}
return json.dumps(value, ensure_ascii=False, default=str)
return document
def _build_payload(
self,
documents: Sequence[str | Document | Mapping[str, Any]],
query: str,
*,
rank_fields: Sequence[str] | None,
model: str | None,
top_n: int | None,
task: str | None,
) -> dict[str, Any]:
"""Build the request body for the `/rerank` endpoint."""
requested_top_n = top_n if top_n is None or top_n > 0 else self.top_n
payload: dict[str, Any] = {
"model": model or self.model,
"query": query,
"documents": [
self._document_to_str(document, rank_fields) for document in documents
],
"return_documents": False,
}
if requested_top_n is not None:
payload["top_n"] = requested_top_n
if task is not None:
payload["task"] = task
return payload
@staticmethod
def _parse_response(response: Mapping[str, Any]) -> list[dict[str, Any]]:
"""Extract index and score pairs from a reranking response."""
return [
{
"index": result["index"],
"relevance_score": result["relevance_score"],
}
for result in response["data"]
]
def rerank(
self,
documents: Sequence[str | Document | Mapping[str, Any]],
query: str,
*,
rank_fields: Sequence[str] | None = None,
model: str | None = None,
top_n: int | None = -1,
task: str | None = None,
) -> list[dict[str, Any]]:
"""Return document indexes ordered by relevance to a query.
Args:
documents: Documents to rerank.
query: Query used for reranking.
rank_fields: Mapping fields to include when serializing mappings.
model: Model to use instead of the configured model.
top_n: Number of results to return. `None` returns all results.
task: Optional task instruction for the reranking model.
Returns:
Reranking results containing each document index and score.
"""
if not documents:
return []
payload = self._build_payload(
documents,
query,
rank_fields=rank_fields,
model=model,
top_n=top_n,
task=task,
)
response = self.client.post("/rerank", cast_to=_RESPONSE_TYPE, body=payload)
return self._parse_response(response)
async def arerank(
self,
documents: Sequence[str | Document | Mapping[str, Any]],
query: str,
*,
rank_fields: Sequence[str] | None = None,
model: str | None = None,
top_n: int | None = -1,
task: str | None = None,
) -> list[dict[str, Any]]:
"""Asynchronously return document indexes ordered by relevance to a query.
Args:
documents: Documents to rerank.
query: Query used for reranking.
rank_fields: Mapping fields to include when serializing mappings.
model: Model to use instead of the configured model.
top_n: Number of results to return. `None` returns all results.
task: Optional task instruction for the reranking model.
Returns:
Reranking results containing each document index and score.
"""
if not documents:
return []
payload = self._build_payload(
documents,
query,
rank_fields=rank_fields,
model=model,
top_n=top_n,
task=task,
)
response = await self.async_client.post(
"/rerank", cast_to=_RESPONSE_TYPE, body=payload
)
return self._parse_response(response)
@staticmethod
def _apply_results(
documents: Sequence[Document],
results: Sequence[Mapping[str, Any]],
) -> list[Document]:
"""Copy the reranked documents, recording each relevance score."""
compressed = []
for result in results:
document = documents[result["index"]]
document_copy = Document(
document.page_content,
metadata=deepcopy(document.metadata),
)
document_copy.metadata["relevance_score"] = result["relevance_score"]
compressed.append(document_copy)
return compressed
@override
def compress_documents(
self,
documents: Sequence[Document],
query: str,
callbacks: Callbacks | None = None,
) -> Sequence[Document]:
"""Compress documents by keeping the most relevant results."""
return self._apply_results(documents, self.rerank(documents, query))
@override
async def acompress_documents(
self,
documents: Sequence[Document],
query: str,
callbacks: Callbacks | None = None,
) -> Sequence[Document]:
"""Asynchronously compress documents by keeping the most relevant results."""
return self._apply_results(documents, await self.arerank(documents, query))
@@ -5,6 +5,7 @@ EXPECTED_ALL = [
"ChatFireworks",
"Fireworks",
"FireworksEmbeddings",
"FireworksRerank",
]
@@ -0,0 +1,217 @@
from typing import Any
import pytest
from langchain_core.documents import Document
from langchain_fireworks import FireworksRerank
class FakeClient:
def __init__(self, response: dict[str, Any]) -> None:
self.response = response
self.calls: list[dict[str, Any]] = []
def post(
self,
path: str,
*,
cast_to: type[dict[str, Any]],
body: dict[str, Any],
) -> dict[str, Any]:
self.calls.append({"path": path, "cast_to": cast_to, "body": body})
return self.response
class FakeAsyncClient(FakeClient):
async def post( # type: ignore[override]
self,
path: str,
*,
cast_to: type[dict[str, Any]],
body: dict[str, Any],
) -> dict[str, Any]:
return super().post(path, cast_to=cast_to, body=body)
def _reranker(
response: dict[str, Any], **kwargs: Any
) -> tuple[FireworksRerank, FakeClient, FakeAsyncClient]:
client = FakeClient(response)
async_client = FakeAsyncClient(response)
reranker = FireworksRerank(client=client, async_client=async_client, **kwargs)
return reranker, client, async_client
def test_missing_api_key_without_clients() -> None:
with pytest.raises(ValueError, match="FIREWORKS_API_KEY is required"):
FireworksRerank(model="reranker", fireworks_api_key=None)
def test_missing_api_key_with_only_sync_client() -> None:
with pytest.raises(ValueError, match="FIREWORKS_API_KEY is required"):
FireworksRerank(
model="reranker",
client=FakeClient({"data": []}),
fireworks_api_key=None,
)
def test_rerank_posts_fireworks_payload() -> None:
reranker, client, _ = _reranker(
{"data": [{"index": 1, "relevance_score": 0.9}]},
model="fireworks/qwen3-reranker-8b",
)
result = reranker.rerank(
[Document("first"), Document("second")],
"the query",
top_n=1,
)
assert result == [{"index": 1, "relevance_score": 0.9}]
assert client.calls == [
{
"path": "/rerank",
"cast_to": dict[str, Any],
"body": {
"model": "fireworks/qwen3-reranker-8b",
"query": "the query",
"documents": ["first", "second"],
"return_documents": False,
"top_n": 1,
},
}
]
async def test_arerank_posts_fireworks_payload() -> None:
reranker, client, async_client = _reranker(
{"data": [{"index": 1, "relevance_score": 0.9}]},
model="fireworks/qwen3-reranker-8b",
)
result = await reranker.arerank(
[Document("first"), Document("second")],
"the query",
top_n=1,
)
assert result == [{"index": 1, "relevance_score": 0.9}]
assert async_client.calls == [
{
"path": "/rerank",
"cast_to": dict[str, Any],
"body": {
"model": "fireworks/qwen3-reranker-8b",
"query": "the query",
"documents": ["first", "second"],
"return_documents": False,
"top_n": 1,
},
}
]
assert client.calls == []
def test_rerank_serializes_mappings_and_rank_fields() -> None:
reranker, client, _ = _reranker({"data": []}, model="reranker")
reranker.rerank(
[{"title": "keep", "body": "keep", "ignored": "drop"}],
"query",
rank_fields=["title", "body"],
top_n=None,
)
assert client.calls[0]["body"] == {
"model": "reranker",
"query": "query",
"documents": ['{"title": "keep", "body": "keep"}'],
"return_documents": False,
}
async def test_arerank_serializes_mappings_and_rank_fields() -> None:
reranker, _, async_client = _reranker({"data": []}, model="reranker")
await reranker.arerank(
[{"title": "keep", "body": "keep", "ignored": "drop"}],
"query",
rank_fields=["title", "body"],
top_n=None,
)
assert async_client.calls[0]["body"] == {
"model": "reranker",
"query": "query",
"documents": ['{"title": "keep", "body": "keep"}'],
"return_documents": False,
}
def test_rerank_empty_documents_does_not_call_client() -> None:
reranker, client, _ = _reranker({"data": []}, model="reranker")
assert reranker.rerank([], "query") == []
assert client.calls == []
async def test_arerank_empty_documents_does_not_call_client() -> None:
reranker, _, async_client = _reranker({"data": []}, model="reranker")
assert await reranker.arerank([], "query") == []
assert async_client.calls == []
def test_compress_documents_preserves_metadata_and_adds_score() -> None:
reranker, _, _ = _reranker(
{
"data": [
{"index": 1, "relevance_score": 0.8},
{"index": 0, "relevance_score": 0.4},
]
},
model="reranker",
)
documents = [
Document("first", metadata={"nested": {"value": 1}}),
Document("second", metadata={"source": "test"}),
]
result = reranker.compress_documents(documents, "query")
assert [document.page_content for document in result] == ["second", "first"]
assert result[0].metadata == {"source": "test", "relevance_score": 0.8}
assert result[1].metadata == {
"nested": {"value": 1},
"relevance_score": 0.4,
}
async def test_acompress_documents_preserves_metadata_and_adds_score() -> None:
reranker, client, async_client = _reranker(
{
"data": [
{"index": 1, "relevance_score": 0.8},
{"index": 0, "relevance_score": 0.4},
]
},
model="reranker",
)
documents = [
Document("first", metadata={"nested": {"value": 1}}),
Document("second", metadata={"source": "test"}),
]
result = await reranker.acompress_documents(documents, "query")
assert [document.page_content for document in result] == ["second", "first"]
assert result[0].metadata == {"source": "test", "relevance_score": 0.8}
assert result[1].metadata == {
"nested": {"value": 1},
"relevance_score": 0.4,
}
assert documents[1].metadata == {"source": "test"}
# The async path must not fall back to the sync client via `run_in_executor`.
assert client.calls == []
assert len(async_client.calls) == 1