diff --git a/libs/partners/fireworks/langchain_fireworks/__init__.py b/libs/partners/fireworks/langchain_fireworks/__init__.py index a8608eb0ca..47159e1ea7 100644 --- a/libs/partners/fireworks/langchain_fireworks/__init__.py +++ b/libs/partners/fireworks/langchain_fireworks/__init__.py @@ -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__", ] diff --git a/libs/partners/fireworks/langchain_fireworks/rerank.py b/libs/partners/fireworks/langchain_fireworks/rerank.py new file mode 100644 index 0000000000..6d21f21ab9 --- /dev/null +++ b/libs/partners/fireworks/langchain_fireworks/rerank.py @@ -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)) diff --git a/libs/partners/fireworks/tests/unit_tests/test_imports.py b/libs/partners/fireworks/tests/unit_tests/test_imports.py index 6b7a77b6ac..ab49a11c3b 100644 --- a/libs/partners/fireworks/tests/unit_tests/test_imports.py +++ b/libs/partners/fireworks/tests/unit_tests/test_imports.py @@ -5,6 +5,7 @@ EXPECTED_ALL = [ "ChatFireworks", "Fireworks", "FireworksEmbeddings", + "FireworksRerank", ] diff --git a/libs/partners/fireworks/tests/unit_tests/test_rerank.py b/libs/partners/fireworks/tests/unit_tests/test_rerank.py new file mode 100644 index 0000000000..a9afbb4fba --- /dev/null +++ b/libs/partners/fireworks/tests/unit_tests/test_rerank.py @@ -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