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
18 changes: 17 additions & 1 deletion graphiti_core/edges.py
Original file line number Diff line number Diff line change
Expand Up @@ -28,7 +28,7 @@
from graphiti_core.driver.driver import GraphDriver, GraphProvider
from graphiti_core.embedder import EmbedderClient
from graphiti_core.errors import EdgeNotFoundError, GroupsEdgesNotFoundError
from graphiti_core.helpers import parse_db_date
from graphiti_core.helpers import parse_db_date, validate_group_id
from graphiti_core.models.edges.edge_db_queries import (
COMMUNITY_EDGE_RETURN,
EPISODIC_EDGE_RETURN,
Expand All @@ -53,6 +53,12 @@ class Edge(BaseModel, ABC):
target_node_uuid: str
created_at: datetime

def _validate_for_write(self) -> None:
# Validate group_id at the persistence boundary only. Hydration from the DB must
# stay tolerant of any stored value, so validation lives here (called by every
# concrete save()) rather than on the model where it would also fire on reads.
validate_group_id(self.group_id)

@abstractmethod
async def save(self, driver: GraphDriver): ...

Expand Down Expand Up @@ -142,6 +148,8 @@ async def get_by_uuid(cls, driver: GraphDriver, uuid: str): ...

class EpisodicEdge(Edge):
async def save(self, driver: GraphDriver):
self._validate_for_write()

if driver.graph_operations_interface:
try:
return await driver.graph_operations_interface.episodic_edge_save(self, driver)
Expand Down Expand Up @@ -333,6 +341,8 @@ async def load_fact_embedding(self, driver: GraphDriver):
self.fact_embedding = records[0]['fact_embedding']

async def save(self, driver: GraphDriver):
self._validate_for_write()

if driver.graph_operations_interface:
try:
return await driver.graph_operations_interface.edge_save(self, driver)
Expand Down Expand Up @@ -574,6 +584,8 @@ async def get_by_node_uuid(cls, driver: GraphDriver, node_uuid: str):

class CommunityEdge(Edge):
async def save(self, driver: GraphDriver):
self._validate_for_write()

if driver.graph_operations_interface:
try:
return await driver.graph_operations_interface.community_edge_save(self, driver)
Expand Down Expand Up @@ -688,6 +700,8 @@ async def get_by_group_ids(

class HasEpisodeEdge(Edge):
async def save(self, driver: GraphDriver):
self._validate_for_write()

if driver.graph_operations_interface:
try:
return await driver.graph_operations_interface.has_episode_edge_save(self, driver)
Expand Down Expand Up @@ -821,6 +835,8 @@ async def get_by_group_ids(

class NextEpisodeEdge(Edge):
async def save(self, driver: GraphDriver):
self._validate_for_write()

if driver.graph_operations_interface:
try:
return await driver.graph_operations_interface.next_episode_edge_save(self, driver)
Expand Down
13 changes: 12 additions & 1 deletion graphiti_core/errors.py
Original file line number Diff line number Diff line change
Expand Up @@ -75,7 +75,7 @@ def __init__(self, entity_type: str, entity_type_attribute: str):
super().__init__(self.message)


class GroupIdValidationError(GraphitiError):
class GroupIdValidationError(GraphitiError, ValueError):
"""Raised when a group_id contains invalid characters."""

def __init__(self, group_id: str):
Expand All @@ -93,3 +93,14 @@ def __init__(self, node_labels: list[str]):
f'alphanumeric characters or underscores: {label_list}'
)
super().__init__(self.message)


class PropertyNameValidationError(GraphitiError, ValueError):
"""Raised when a property filter name is not a safe Cypher identifier."""

def __init__(self, property_name: str):
self.message = (
f'property filter name "{property_name}" must start with a letter or underscore '
'and contain only alphanumeric characters or underscores'
)
super().__init__(self.message)
20 changes: 19 additions & 1 deletion graphiti_core/helpers.py
Original file line number Diff line number Diff line change
Expand Up @@ -28,7 +28,11 @@
from pydantic import BaseModel

from graphiti_core.driver.driver import GraphProvider
from graphiti_core.errors import GroupIdValidationError, NodeLabelValidationError
from graphiti_core.errors import (
GroupIdValidationError,
NodeLabelValidationError,
PropertyNameValidationError,
)

load_dotenv()

Expand Down Expand Up @@ -186,6 +190,20 @@ def validate_node_labels(node_labels: list[str] | None) -> bool:
return True


def validate_property_name(property_name: str) -> bool:
"""Validate that a property name is safe to interpolate into Cypher expressions.

Property keys cannot be parameterized in Cypher, so filter property names are
interpolated directly into the query and must be constrained to a safe identifier
pattern to prevent query injection.
"""

if not SAFE_CYPHER_IDENTIFIER_PATTERN.match(property_name or ''):
raise PropertyNameValidationError(property_name)

return True


def validate_excluded_entity_types(
excluded_entity_types: list[str] | None, entity_types: dict[str, type[BaseModel]] | None = None
) -> bool:
Expand Down
16 changes: 15 additions & 1 deletion graphiti_core/nodes.py
Original file line number Diff line number Diff line change
Expand Up @@ -32,7 +32,7 @@
)
from graphiti_core.embedder import EmbedderClient
from graphiti_core.errors import NodeNotFoundError
from graphiti_core.helpers import parse_db_date, validate_node_labels
from graphiti_core.helpers import parse_db_date, validate_group_id, validate_node_labels
from graphiti_core.models.nodes.node_db_queries import (
COMMUNITY_NODE_RETURN,
COMMUNITY_NODE_RETURN_NEPTUNE,
Expand Down Expand Up @@ -105,6 +105,12 @@ def validate_labels(cls, value: list[str]) -> list[str]:
validate_node_labels(value)
return value

def _validate_for_write(self) -> None:
# Validate group_id at the persistence boundary only. Hydration from the DB must
# stay tolerant of any stored value, so validation lives here (called by every
# concrete save()) rather than on the model where it would also fire on reads.
validate_group_id(self.group_id)

@abstractmethod
async def save(self, driver: GraphDriver): ...

Expand Down Expand Up @@ -332,6 +338,8 @@ class EpisodicNode(Node):
)

async def save(self, driver: GraphDriver):
self._validate_for_write()

if driver.graph_operations_interface:
try:
return await driver.graph_operations_interface.episodic_node_save(self, driver)
Expand Down Expand Up @@ -544,6 +552,8 @@ async def load_name_embedding(self, driver: GraphDriver):
self.name_embedding = records[0]['name_embedding']

async def save(self, driver: GraphDriver):
self._validate_for_write()

if driver.graph_operations_interface:
try:
return await driver.graph_operations_interface.node_save(self, driver)
Expand Down Expand Up @@ -689,6 +699,8 @@ class CommunityNode(Node):
summary: str = Field(description='region summary of member nodes', default_factory=str)

async def save(self, driver: GraphDriver):
self._validate_for_write()

if driver.graph_operations_interface:
try:
return await driver.graph_operations_interface.community_node_save(self, driver)
Expand Down Expand Up @@ -876,6 +888,8 @@ class SagaNode(Node):
last_summarized_episode_valid_at: datetime | None = None

async def save(self, driver: GraphDriver):
self._validate_for_write()

if driver.graph_operations_interface:
try:
return await driver.graph_operations_interface.saga_node_save(self, driver)
Expand Down
58 changes: 57 additions & 1 deletion graphiti_core/search/search_filters.py
Original file line number Diff line number Diff line change
Expand Up @@ -21,7 +21,7 @@
from pydantic import BaseModel, Field, field_validator

from graphiti_core.driver.driver import GraphProvider
from graphiti_core.helpers import validate_node_labels
from graphiti_core.helpers import validate_node_labels, validate_property_name


class ComparisonOperator(Enum):
Expand Down Expand Up @@ -64,6 +64,9 @@ class SearchFilters(BaseModel):
created_at: list[list[DateFilter]] | None = Field(default=None)
expired_at: list[list[DateFilter]] | None = Field(default=None)
edge_uuids: list[str] | None = Field(default=None)
# A single property_filters list is intentionally shared: it is applied to the node
# alias `n` by node_search_filter_query_constructor and to the edge alias `e` by
# edge_search_filter_query_constructor.
property_filters: list[PropertyFilter] | None = Field(default=None)

@field_validator('node_labels')
Expand All @@ -72,6 +75,16 @@ def validate_node_label_filters(cls, value: list[str] | None) -> list[str] | Non
validate_node_labels(value)
return value

@field_validator('property_filters')
@classmethod
def validate_property_filter_names(
cls, value: list[PropertyFilter] | None
) -> list[PropertyFilter] | None:
if value is not None:
for property_filter in value:
validate_property_name(property_filter.property_name)
return value


def cypher_to_opensearch_operator(op: ComparisonOperator) -> str:
mapping = {
Expand All @@ -83,6 +96,35 @@ def cypher_to_opensearch_operator(op: ComparisonOperator) -> str:
return mapping.get(op, op.value)


def property_filter_query_constructor(
entity_alias: str,
property_filters: list[PropertyFilter],
param_prefix: str,
) -> tuple[list[str], dict[str, Any]]:
"""Build Cypher fragments for a list of property filters against a single entity alias.

Property keys cannot be parameterized in Cypher, so each name is validated
(defense-in-depth against model_construct()/other validation bypasses) before it is
interpolated. Property values are always passed as query parameters.
"""
filter_queries: list[str] = []
filter_params: dict[str, Any] = {}

for i, property_filter in enumerate(property_filters):
validate_property_name(property_filter.property_name)
property_reference = f'{entity_alias}.{property_filter.property_name}'
operator = property_filter.comparison_operator

if operator in (ComparisonOperator.is_null, ComparisonOperator.is_not_null):
filter_queries.append(f'{property_reference} {operator.value}')
else:
param_name = f'{param_prefix}_{i}'
filter_queries.append(f'{property_reference} {operator.value} ${param_name}')
filter_params[param_name] = property_filter.property_value

return filter_queries, filter_params


def node_search_filter_query_constructor(
filters: SearchFilters,
provider: GraphProvider,
Expand All @@ -101,6 +143,13 @@ def node_search_filter_query_constructor(
node_label_filter = 'n:' + node_labels
filter_queries.append(node_label_filter)

if filters.property_filters is not None:
property_queries, property_params = property_filter_query_constructor(
'n', filters.property_filters, 'node_prop'
)
filter_queries.extend(property_queries)
filter_params.update(property_params)

return filter_queries, filter_params


Expand Down Expand Up @@ -133,6 +182,13 @@ def edge_search_filter_query_constructor(
filter_queries.append('e.uuid in $edge_uuids')
filter_params['edge_uuids'] = filters.edge_uuids

if filters.property_filters is not None:
property_queries, property_params = property_filter_query_constructor(
'e', filters.property_filters, 'edge_prop'
)
filter_queries.extend(property_queries)
filter_params.update(property_params)

if filters.node_labels is not None:
# Defense-in-depth for model_construct()/other validation bypasses.
validate_node_labels(filters.node_labels)
Expand Down
8 changes: 7 additions & 1 deletion graphiti_core/utils/bulk_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -31,7 +31,7 @@
from graphiti_core.edges import Edge, EntityEdge, EpisodicEdge, create_entity_edge_embeddings
from graphiti_core.embedder import EmbedderClient
from graphiti_core.graphiti_types import GraphitiClients
from graphiti_core.helpers import normalize_l2, semaphore_gather
from graphiti_core.helpers import normalize_l2, semaphore_gather, validate_group_id
from graphiti_core.models.edges.edge_db_queries import (
get_entity_edge_save_bulk_query,
get_episodic_edge_save_bulk_query,
Expand Down Expand Up @@ -157,6 +157,12 @@ async def add_nodes_and_edges_bulk_tx(
embedder: EmbedderClient,
driver: GraphDriver,
):
# Validate group_ids on write only. Hydration from the DB stays tolerant of any
# stored value, so the bulk persistence boundary re-validates here (the per-object
# save() methods are bypassed on this path).
for element in (*episodic_nodes, *episodic_edges, *entity_nodes, *entity_edges):
validate_group_id(element.group_id)

episodes = [dict(episode) for episode in episodic_nodes]
for episode in episodes:
episode['source'] = str(episode['source'].value)
Expand Down
26 changes: 25 additions & 1 deletion server/graph_service/dto/ingest.py
Original file line number Diff line number Diff line change
@@ -1,15 +1,39 @@
from pydantic import BaseModel, Field
from graphiti_core.errors import GroupIdValidationError # type: ignore
from graphiti_core.helpers import validate_group_id # type: ignore
from pydantic import BaseModel, Field, field_validator

from graph_service.dto.common import Message


def _validate_request_group_id(value: str) -> str:
# Re-raise as ValueError so Pydantic wraps it into a ValidationError (HTTP 422)
# regardless of whether the installed graphiti-core makes GroupIdValidationError a
# ValueError subclass. Rejecting bad group_ids on write keeps records reachable and
# deletable via the read/delete API paths, which validate the same pattern.
try:
validate_group_id(value)
except GroupIdValidationError as error:
raise ValueError(str(error)) from error
return value


class AddMessagesRequest(BaseModel):
group_id: str = Field(..., description='The group id of the messages to add')
messages: list[Message] = Field(..., description='The messages to add')

@field_validator('group_id')
@classmethod
def validate_group_id_field(cls, value: str) -> str:
return _validate_request_group_id(value)


class AddEntityNodeRequest(BaseModel):
uuid: str = Field(..., description='The uuid of the node to add')
group_id: str = Field(..., description='The group id of the node to add')
name: str = Field(..., description='The name of the node to add')
summary: str = Field(default='', description='The summary of the node to add')

@field_validator('group_id')
@classmethod
def validate_group_id_field(cls, value: str) -> str:
return _validate_request_group_id(value)
29 changes: 29 additions & 0 deletions server/tests/test_dto_validation.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,29 @@
"""Unit tests for request DTO validation (no database or network required)."""

import pytest
from pydantic import ValidationError

from graph_service.dto import AddEntityNodeRequest, AddMessagesRequest


def test_add_entity_node_request_rejects_unsafe_group_id():
# The write path must reject group_ids the read/delete path would refuse to match,
# so a bad group_id surfaces as a 422 (ValidationError) rather than creating an
# unreachable record.
with pytest.raises(ValidationError):
AddEntityNodeRequest(uuid='u', group_id='bad"group', name='n')


def test_add_entity_node_request_accepts_valid_group_id():
request = AddEntityNodeRequest(uuid='u', group_id='valid-group_1', name='n')
assert request.group_id == 'valid-group_1'


def test_add_messages_request_rejects_unsafe_group_id():
with pytest.raises(ValidationError):
AddMessagesRequest(group_id='bad"group', messages=[])


def test_add_messages_request_accepts_valid_group_id():
request = AddMessagesRequest(group_id='valid-group_1', messages=[])
assert request.group_id == 'valid-group_1'
Loading
Loading