Skip to content
Open
Show file tree
Hide file tree
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
1 change: 1 addition & 0 deletions CHANGELOG.md
Original file line number Diff line number Diff line change
Expand Up @@ -5,6 +5,7 @@
### 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

Expand Down
21 changes: 2 additions & 19 deletions src/neo4j_graphrag/llm/anthropic_llm.py
Original file line number Diff line number Diff line change
Expand Up @@ -14,7 +14,6 @@
from __future__ import annotations

import json
import warnings
from typing import (
TYPE_CHECKING,
Any,
Expand All @@ -26,12 +25,6 @@
cast,
)

# 3rd party dependencies
try:
import httpx
except ImportError:
httpx = None # type: ignore[assignment]

from pydantic import BaseModel, ValidationError

from neo4j_graphrag.exceptions import LLMGenerationError
Expand All @@ -43,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 (
Expand Down Expand Up @@ -225,18 +219,7 @@ def __init__(
**kwargs,
)
self.anthropic = anthropic
http_client = kwargs.pop("http_client", None)
sync_params = kwargs.copy()
async_params = kwargs.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 = anthropic.Anthropic(**sync_params)
self.async_client = anthropic.AsyncAnthropic(**async_params)

Expand Down
34 changes: 3 additions & 31 deletions src/neo4j_graphrag/llm/openai_llm.py
Original file line number Diff line number Diff line change
Expand Up @@ -19,7 +19,6 @@
import abc
import json
import logging
import warnings
from typing import (
TYPE_CHECKING,
Any,
Expand All @@ -34,10 +33,6 @@
)

# 3rd party dependencies
try:
import httpx
except ImportError:
httpx = None # type: ignore[assignment]
from pydantic import BaseModel, ValidationError

# project dependencies
Expand Down Expand Up @@ -66,6 +61,7 @@
ToolCallResponse,
UserMessage,
)
from .utils import split_http_client_kwargs

if TYPE_CHECKING:
from openai import AsyncOpenAI, OpenAI
Expand Down Expand Up @@ -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)

Expand Down Expand Up @@ -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)
52 changes: 51 additions & 1 deletion src/neo4j_graphrag/llm/utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -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."""
Expand Down Expand Up @@ -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
57 changes: 57 additions & 0 deletions tests/unit/llm/test_llm_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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()