Skip to content
Closed
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
4 changes: 4 additions & 0 deletions CHANGELOG.md
Original file line number Diff line number Diff line change
Expand Up @@ -2,6 +2,10 @@

## Next

### Added

- Added `BaseVertexAILLM`, a new base class holding all of `VertexAILLM`'s shared message-building, generation-config/schema-handling, and response-parsing logic. Unlike `BaseAnthropicLLM`/`BaseOpenAILLM` (which hold a persistent SDK client for subclasses to construct), `VertexAILLM` has no per-instance client — it relies on a global `vertexai.init(...)` plus a fresh `GenerativeModel` per call — so `BaseVertexAILLM` instead declares a single abstract `_get_model(...)` hook, which `VertexAILLM` implements exactly as before. `BaseVertexAILLM` is exported from `neo4j_graphrag.llm` as a documented extension point.

## 1.19.0

### Added
Expand Down
7 changes: 7 additions & 0 deletions docs/source/api.rst
Original file line number Diff line number Diff line change
Expand Up @@ -388,6 +388,13 @@ OllamaLLM
:members:


BaseVertexAILLM
---------------

.. autoclass:: neo4j_graphrag.llm.vertexai_llm.BaseVertexAILLM
:members:


VertexAILLM
-----------

Expand Down
3 changes: 2 additions & 1 deletion src/neo4j_graphrag/llm/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -31,14 +31,15 @@
from .openai_llm import AzureOpenAILLM, BaseOpenAILLM, OpenAILLM
from .types import LLMResponse, LLMUsage
from .utils import split_http_client_kwargs
from .vertexai_llm import VertexAILLM
from .vertexai_llm import BaseVertexAILLM, VertexAILLM

__all__ = [
"GEMINI_DEFAULT_IMAGE_MIME_TYPE",
"GEMINI_SUPPORTED_IMAGE_MIME_TYPES",
"AnthropicLLM",
"BaseAnthropicLLM",
"BaseGeminiLLM",
"BaseVertexAILLM",
"BedrockLLM",
"CohereLLM",
"GeminiImageMimeType",
Expand Down
66 changes: 58 additions & 8 deletions src/neo4j_graphrag/llm/vertexai_llm.py
Original file line number Diff line number Diff line change
Expand Up @@ -15,6 +15,7 @@
# built-in dependencies
from __future__ import annotations

import abc
import inspect
import logging
from typing import Any, List, Optional, Sequence, Type, Union, cast
Expand Down Expand Up @@ -113,8 +114,8 @@ def _extract_generation_config_params(


# pylint: disable=arguments-differ, redefined-builtin, no-else-return
class VertexAILLM(LLMBase):
"""Interface for large language models on Vertex AI
class BaseVertexAILLM(LLMBase, abc.ABC):
"""Base class for Vertex AI LLMs.

Args:
model_name (str, optional): Name of the LLM to use. Defaults to "gemini-2.5-flash".
Expand All @@ -128,6 +129,14 @@ class VertexAILLM(LLMBase):
Raises:
LLMGenerationError: If there's an error generating the response from the model.

Holds all the shared message-building, generation-config/schema-handling,
and response-parsing logic. Unlike the Anthropic/OpenAI/Gemini base
classes, there is no persistent per-instance SDK client to construct here
(Vertex AI relies on a global ``vertexai.init(...)`` plus a fresh
``GenerativeModel`` per call) — subclasses are only responsible for
implementing :meth:`_get_model`, which controls how that model object is
constructed.

Example:

.. code-block:: python
Expand Down Expand Up @@ -401,16 +410,19 @@ def _get_llm_tools(
)
]

@abc.abstractmethod
def _get_model(
self,
system_instruction: Optional[str] = None,
) -> GenerativeModel:
# system_message = [system_instruction] if system_instruction is not None else []
model = GenerativeModel(
model_name=self.model_name,
system_instruction=system_instruction,
)
return model
"""Construct the ``GenerativeModel`` used for a single call.

This is the one thing a subclass is responsible for — everything
else (input building, generation-config/schema handling, response
parsing) is inherited unchanged. A subclass reaching a different
endpoint or hosting configuration (e.g. via ``vertexai.init(...)``
called with different arguments) only needs to override this method.
"""

def get_messages(
self,
Expand Down Expand Up @@ -601,3 +613,41 @@ def _parse_content_response(self, response: GenerationResponse) -> LLMResponse:
total_tokens=metadata.total_token_count,
)
return LLMResponse(content=response.text, usage=usage)


class VertexAILLM(BaseVertexAILLM):
"""Interface for large language models on Vertex AI

Args:
model_name (str, optional): Name of the LLM to use. Defaults to "gemini-2.5-flash".
model_params (Optional[dict], optional): Additional parameters for LLMInterface(V1) passed to the model when text is sent to it. Defaults to None.
system_instruction: Optional[str], optional): Additional instructions for setting the behavior and context for the model in a conversation. Defaults to None.
rate_limit_handler (Optional[RateLimitHandler], optional): Rate limit handler for LLMInterface(V1). Defaults to None.
**kwargs (Any): Arguments passed to the model when for the class is initialised. Defaults to None.

Raises:
LLMGenerationError: If there's an error generating the response from the model.

Example:

.. code-block:: python

from neo4j_graphrag.llm import VertexAILLM
from vertexai.generative_models import GenerationConfig

generation_config = GenerationConfig(temperature=0.0)
llm = VertexAILLM(
model_name="gemini-2.5-flash", generation_config=generation_config
)
llm.invoke("Who is the mother of Paul Atreides?")
"""

def _get_model(
self,
system_instruction: Optional[str] = None,
) -> GenerativeModel:
model = GenerativeModel(
model_name=self.model_name,
system_instruction=system_instruction,
)
return model
59 changes: 57 additions & 2 deletions tests/unit/llm/test_vertexai_llm.py
Original file line number Diff line number Diff line change
Expand Up @@ -13,20 +13,21 @@
# limitations under the License.
from __future__ import annotations

from typing import cast
from typing import Optional, cast
from typing import List

from unittest.mock import AsyncMock, MagicMock, Mock, patch
import pytest
from vertexai.generative_models import (
Content,
GenerationResponse,
GenerativeModel,
Part,
)

from neo4j_graphrag.exceptions import LLMGenerationError
from neo4j_graphrag.llm.types import ToolCallResponse
from neo4j_graphrag.llm.vertexai_llm import VertexAILLM
from neo4j_graphrag.llm.vertexai_llm import BaseVertexAILLM, VertexAILLM
from neo4j_graphrag.tool import Tool
from neo4j_graphrag.types import LLMMessage
from neo4j_graphrag.utils.rate_limit import NoOpRateLimitHandler
Expand Down Expand Up @@ -660,3 +661,57 @@ def test_vertexai_llm_default_model_matches_docstring(
llm = VertexAILLM()
assert llm.model_name == "gemini-2.5-flash"
assert 'Defaults to "gemini-2.5-flash"' in (VertexAILLM.__doc__ or "")


def test_vertexai_llm_is_base_vertexai_llm_subclass() -> None:
assert issubclass(VertexAILLM, BaseVertexAILLM)


def test_base_vertexai_llm_cannot_be_instantiated_without_get_model() -> None:
"""BaseVertexAILLM declares _get_model as abstract; unlike the other base
classes (which have no remaining abstract methods once invoke/ainvoke are
implemented), this one genuinely cannot be instantiated directly."""
with pytest.raises(TypeError):
BaseVertexAILLM(model_name="gemini-1.5-flash-001") # type: ignore[abstract]


@patch("neo4j_graphrag.llm.vertexai_llm.GenerativeModel")
def test_vertexai_llm_get_model_is_the_only_override(
GenerativeModelMock: MagicMock,
) -> None:
"""VertexAILLM's one responsibility is implementing _get_model; its
GenerativeModel construction is unchanged from before the extraction."""
llm = VertexAILLM(model_name="gemini-1.5-flash-001")

assert VertexAILLM._get_model is not BaseVertexAILLM._get_model

model = llm._get_model(system_instruction="be nice")
GenerativeModelMock.assert_called_once_with(
model_name="gemini-1.5-flash-001", system_instruction="be nice"
)
assert model is GenerativeModelMock.return_value


def test_minimal_base_vertexai_llm_subclass_exercises_invoke() -> None:
"""The exported extension contract: a subclass that only implements
_get_model should get message building, generation-config handling, and
response parsing from BaseVertexAILLM, and its custom model must be the
one invoke() actually calls."""
custom_model = MagicMock()
mock_response = Mock()
mock_response.text = "custom model response"
mock_response.usage_metadata = None
custom_model.generate_content.return_value = mock_response

class MinimalVertexAILLM(BaseVertexAILLM):
def _get_model(
self,
system_instruction: Optional[str] = None,
) -> GenerativeModel:
return cast(GenerativeModel, custom_model)

llm = MinimalVertexAILLM(model_name="gemini-1.5-flash-001")
response = llm.invoke("hello")

assert response.content == "custom model response"
custom_model.generate_content.assert_called_once()
Loading