mirror of
https://github.com/langchain-ai/langchain.git
synced 2026-10-05 17:35:28 +03:00
feat(fireworks): add document reranking (#39732)
This commit is contained in:
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
|
||||
Reference in new issue
Block a user