diff --git a/CHANGELOG.md b/CHANGELOG.md index e54fff1dd..f89448b73 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -5,11 +5,16 @@ ### Added - `AnthropicLLM` now supports structured output via the `response_format` argument, accepting a Pydantic model or an Anthropic `output_config` dict, alongside `OpenAILLM` and `VertexAILLM`. +- Added `neo4j_graphrag.llm.utils.split_http_client_kwargs`, a shared helper that routes a constructor's `http_client` kwarg to whichever of the sync/async SDK clients it matches. `AnthropicLLM`, `OpenAILLM`, and `AzureOpenAILLM` now all use this single implementation instead of three separately maintained copies of the same logic. Custom subclasses that construct their own SDK clients can call it to get the same behavior. ### Changed - (**breaking**) `AnthropicLLM.supports_structured_output` is now `True`. As a result, `SchemaFromTextExtractor` and `LLMEntityRelationExtractor` (and `SimpleKGPipeline`, which enables structured output automatically when the LLM supports it) now use structured output by default with `AnthropicLLM`. This requires a Claude 4.5+ model (e.g. `claude-sonnet-4-5`); using `AnthropicLLM` with an older Claude model in these components will now raise an error where it previously worked. To keep the previous behavior, use a Claude 4.5+ model, or construct `LLMEntityRelationExtractor` / `SchemaFromTextExtractor` directly with `use_structured_output=False`. +### Fixed + +- Fixed a bug in `AnthropicLLM` where an `http_client` passed via kwargs (whether an `httpx.Client` or `httpx.AsyncClient`) was forwarded to both the sync `anthropic.Anthropic` and async `anthropic.AsyncAnthropic` clients, causing a type mismatch. `http_client` is now routed to the matching sync/async client only; other kwargs remain shared. An `http_client` of an unrecognized type now emits a warning and is ignored instead of raising, matching `OpenAILLM`'s existing behavior. + ## 1.18.0 ### Changed diff --git a/src/neo4j_graphrag/llm/anthropic_llm.py b/src/neo4j_graphrag/llm/anthropic_llm.py index 7774f0be1..cf8c62ba3 100644 --- a/src/neo4j_graphrag/llm/anthropic_llm.py +++ b/src/neo4j_graphrag/llm/anthropic_llm.py @@ -36,6 +36,7 @@ MessageList, UserMessage, ) +from neo4j_graphrag.llm.utils import split_http_client_kwargs from neo4j_graphrag.message_history import MessageHistory from neo4j_graphrag.types import LLMMessage from neo4j_graphrag.utils.rate_limit import ( @@ -218,8 +219,9 @@ def __init__( **kwargs, ) self.anthropic = anthropic - self.client = anthropic.Anthropic(**kwargs) - self.async_client = anthropic.AsyncAnthropic(**kwargs) + sync_params, async_params = split_http_client_kwargs(kwargs) + self.client = anthropic.Anthropic(**sync_params) + self.async_client = anthropic.AsyncAnthropic(**async_params) def invoke( self, diff --git a/src/neo4j_graphrag/llm/openai_llm.py b/src/neo4j_graphrag/llm/openai_llm.py index b90826ca7..47645d4e3 100644 --- a/src/neo4j_graphrag/llm/openai_llm.py +++ b/src/neo4j_graphrag/llm/openai_llm.py @@ -19,7 +19,6 @@ import abc import json import logging -import warnings from typing import ( TYPE_CHECKING, Any, @@ -34,10 +33,6 @@ ) # 3rd party dependencies -try: - import httpx -except ImportError: - httpx = None # type: ignore[assignment] from pydantic import BaseModel, ValidationError # project dependencies @@ -66,6 +61,7 @@ ToolCallResponse, UserMessage, ) +from .utils import split_http_client_kwargs if TYPE_CHECKING: from openai import AsyncOpenAI, OpenAI @@ -658,19 +654,7 @@ def __init__( model_params=model_params, rate_limit_handler=rate_limit_handler, ) - http_client = kwargs.pop("http_client", None) - params = kwargs.copy() - sync_params = params.copy() - async_params = params.copy() - if httpx is not None and isinstance(http_client, httpx.Client): - sync_params["http_client"] = http_client - elif httpx is not None and isinstance(http_client, httpx.AsyncClient): - async_params["http_client"] = http_client - elif http_client is not None: - warnings.warn( - f"Invalid http_client type (got {type(http_client)}, expected httpx.Client or httpx.AsyncClient). Using default client.", - stacklevel=2, - ) + sync_params, async_params = split_http_client_kwargs(kwargs) self.client = self.openai.OpenAI(**sync_params) self.async_client = self.openai.AsyncOpenAI(**async_params) @@ -700,18 +684,6 @@ def __init__( model_params=model_params, rate_limit_handler=rate_limit_handler, ) - http_client = kwargs.pop("http_client", None) - params = kwargs.copy() - sync_params = params.copy() - async_params = params.copy() - if httpx is not None and isinstance(http_client, httpx.Client): - sync_params["http_client"] = http_client - elif httpx is not None and isinstance(http_client, httpx.AsyncClient): - async_params["http_client"] = http_client - elif http_client is not None: - warnings.warn( - f"Invalid http_client type (got {type(http_client)}, expected httpx.Client or httpx.AsyncClient). Using default client.", - stacklevel=2, - ) + sync_params, async_params = split_http_client_kwargs(kwargs) self.client = self.openai.AzureOpenAI(**sync_params) self.async_client = self.openai.AsyncAzureOpenAI(**async_params) diff --git a/src/neo4j_graphrag/llm/utils.py b/src/neo4j_graphrag/llm/utils.py index c2912d4e3..d2fe9d28b 100644 --- a/src/neo4j_graphrag/llm/utils.py +++ b/src/neo4j_graphrag/llm/utils.py @@ -14,13 +14,18 @@ # limitations under the License. from __future__ import annotations import warnings -from typing import Union, Optional +from typing import Any, Union, Optional from pydantic import TypeAdapter from neo4j_graphrag.message_history import MessageHistory from neo4j_graphrag.types import LLMMessage +try: + import httpx +except ImportError: + httpx = None # type: ignore[assignment] + def system_instruction_from_messages(messages: list[LLMMessage]) -> str | None: """Extracts the system instruction from a list of LLMMessage, if present.""" @@ -70,3 +75,48 @@ def legacy_inputs_to_messages( # prompt is a MessageHistory instance messages.extend(prompt.messages) return messages + + +def split_http_client_kwargs( + kwargs: dict[str, Any], +) -> tuple[dict[str, Any], dict[str, Any]]: + """Splits a shared ``kwargs`` dict into separate sync/async constructor kwargs, + routing an optional ``http_client`` to whichever SDK client it matches the type of. + + Several provider integrations (``AnthropicLLM``, ``OpenAILLM``, ``AzureOpenAILLM``) + build both a sync and an async SDK client from a single constructor ``**kwargs`` + dict. Most kwargs (``api_key``, ``max_retries``, ``default_headers``, ...) are safe + to share as-is, but ``http_client`` is not: the sync client needs an + ``httpx.Client``, the async client needs an ``httpx.AsyncClient``, and passing the + wrong type to either raises or silently misbehaves depending on SDK version. + + This pops ``http_client`` out of *kwargs* and returns two independent copies of the + remaining kwargs, with ``http_client`` added back to only the copy whose SDK client + it matches. If ``http_client`` doesn't match either expected type, a warning is + emitted and it is dropped from both, falling back to each SDK's default transport. + + Args: + kwargs: The shared constructor kwargs, as passed by a caller to e.g. + ``AnthropicLLM(...)``. Not mutated. + + Returns: + A ``(sync_kwargs, async_kwargs)`` tuple, each a shallow copy of *kwargs* minus + ``http_client``, with ``http_client`` reinstated in whichever of the two it + belongs to. + """ + kwargs = dict(kwargs) + http_client = kwargs.pop("http_client", None) + sync_kwargs = kwargs.copy() + async_kwargs = kwargs.copy() + if httpx is not None and isinstance(http_client, httpx.Client): + sync_kwargs["http_client"] = http_client + elif httpx is not None and isinstance(http_client, httpx.AsyncClient): + async_kwargs["http_client"] = http_client + elif http_client is not None: + # stacklevel=3 attributes the warning to the caller of the LLM + # constructor, not to the constructor's own call into this helper. + warnings.warn( + f"Invalid http_client type (got {type(http_client)}, expected httpx.Client or httpx.AsyncClient). Using default client.", + stacklevel=3, + ) + return sync_kwargs, async_kwargs diff --git a/tests/unit/llm/test_anthropic_llm.py b/tests/unit/llm/test_anthropic_llm.py index 113fd3637..6e9f4bae4 100644 --- a/tests/unit/llm/test_anthropic_llm.py +++ b/tests/unit/llm/test_anthropic_llm.py @@ -19,6 +19,7 @@ from unittest.mock import AsyncMock, MagicMock, Mock, patch import anthropic +import httpx import pytest from neo4j_graphrag.exceptions import LLMGenerationError from neo4j_graphrag.experimental.components.types import Neo4jGraph @@ -538,6 +539,66 @@ def test_anthropic_llm_close(mock_anthropic: Mock) -> None: mock_anthropic.AsyncAnthropic.return_value.close.assert_called_once() +# --------------------------------------------------------------------------- +# http_client sync/async routing tests +# --------------------------------------------------------------------------- + + +def test_anthropic_llm_with_httpx_client(mock_anthropic: Mock) -> None: + """Test that httpx.Client is forwarded only to the sync Anthropic client.""" + http_client = httpx.Client() + try: + with warnings.catch_warnings(record=True) as caught: + warnings.simplefilter("always") + AnthropicLLM(model_name="claude-3-opus-20240229", http_client=http_client) + + assert not any("Invalid http_client" in str(w.message) for w in caught) + _, sync_kwargs = mock_anthropic.Anthropic.call_args + assert sync_kwargs.get("http_client") is http_client + _, async_kwargs = mock_anthropic.AsyncAnthropic.call_args + assert "http_client" not in async_kwargs + finally: + http_client.close() + + +@pytest.mark.asyncio +async def test_anthropic_llm_with_httpx_async_client(mock_anthropic: Mock) -> None: + """Test that httpx.AsyncClient is forwarded only to the async Anthropic client.""" + async_http_client = httpx.AsyncClient() + with warnings.catch_warnings(record=True) as caught: + warnings.simplefilter("always") + AnthropicLLM(model_name="claude-3-opus-20240229", http_client=async_http_client) + + assert not any("Invalid http_client" in str(w.message) for w in caught) + _, sync_kwargs = mock_anthropic.Anthropic.call_args + assert "http_client" not in sync_kwargs + _, async_kwargs = mock_anthropic.AsyncAnthropic.call_args + assert async_kwargs.get("http_client") is async_http_client + + await async_http_client.aclose() + + +def test_anthropic_llm_no_http_client_no_warning(mock_anthropic: Mock) -> None: + """Test that omitting http_client does not emit a warning.""" + with warnings.catch_warnings(record=True) as caught: + warnings.simplefilter("always") + AnthropicLLM(model_name="claude-3-opus-20240229") + + assert not any("Invalid http_client" in str(w.message) for w in caught) + + +def test_anthropic_llm_with_invalid_http_client_warns(mock_anthropic: Mock) -> None: + """Test that an invalid http_client type emits a warning and falls back to + default construction for both clients.""" + with pytest.warns(UserWarning, match="Invalid http_client type"): + AnthropicLLM(model_name="claude-3-opus-20240229", http_client="not-a-client") + + _, sync_kwargs = mock_anthropic.Anthropic.call_args + _, async_kwargs = mock_anthropic.AsyncAnthropic.call_args + assert "http_client" not in sync_kwargs + assert "http_client" not in async_kwargs + + @pytest.mark.asyncio async def test_anthropic_llm_aclose(mock_anthropic: Mock) -> None: mock_anthropic.AsyncAnthropic.return_value.close = AsyncMock() diff --git a/tests/unit/llm/test_llm_utils.py b/tests/unit/llm/test_llm_utils.py index a47fb79b3..1eb75126d 100644 --- a/tests/unit/llm/test_llm_utils.py +++ b/tests/unit/llm/test_llm_utils.py @@ -12,10 +12,12 @@ # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. # See the License for the specific language governing permissions and # limitations under the License. +import httpx import pytest from neo4j_graphrag.llm.utils import ( legacy_inputs_to_messages, + split_http_client_kwargs, system_instruction_from_messages, ) from neo4j_graphrag.message_history import InMemoryMessageHistory @@ -63,3 +65,58 @@ def test_legacy_inputs_prompt_as_message_history() -> None: result = legacy_inputs_to_messages(history) assert len(result) == 1 assert result[0]["content"] == "from history" + + +def test_split_http_client_kwargs_no_http_client() -> None: + sync_kwargs, async_kwargs = split_http_client_kwargs({"api_key": "sk-test"}) + assert sync_kwargs == {"api_key": "sk-test"} + assert async_kwargs == {"api_key": "sk-test"} + assert "http_client" not in sync_kwargs + assert "http_client" not in async_kwargs + + +def test_split_http_client_kwargs_routes_sync_client() -> None: + client = httpx.Client() + try: + sync_kwargs, async_kwargs = split_http_client_kwargs( + {"api_key": "sk-test", "http_client": client} + ) + assert sync_kwargs["http_client"] is client + assert sync_kwargs["api_key"] == "sk-test" + assert "http_client" not in async_kwargs + assert async_kwargs["api_key"] == "sk-test" + finally: + client.close() + + +def test_split_http_client_kwargs_routes_async_client() -> None: + client = httpx.AsyncClient() + sync_kwargs, async_kwargs = split_http_client_kwargs( + {"api_key": "sk-test", "http_client": client} + ) + assert async_kwargs["http_client"] is client + assert async_kwargs["api_key"] == "sk-test" + assert "http_client" not in sync_kwargs + assert sync_kwargs["api_key"] == "sk-test" + + +def test_split_http_client_kwargs_invalid_type_warns_and_drops() -> None: + with pytest.warns(UserWarning, match="Invalid http_client type"): + sync_kwargs, async_kwargs = split_http_client_kwargs( + {"api_key": "sk-test", "http_client": object()} + ) + assert "http_client" not in sync_kwargs + assert "http_client" not in async_kwargs + assert sync_kwargs["api_key"] == "sk-test" + assert async_kwargs["api_key"] == "sk-test" + + +def test_split_http_client_kwargs_does_not_mutate_input() -> None: + client = httpx.Client() + try: + original: dict[str, object] = {"api_key": "sk-test", "http_client": client} + original_copy = dict(original) + split_http_client_kwargs(original) + assert original == original_copy + finally: + client.close()