import asyncio
import inspect
import json
import textwrap
import sys
import time
from datetime import datetime, timezone
from types import MethodType
import main
import graphiti_core.graphiti as graphiti_module
def instrument_async(name, phase):
original = getattr(graphiti_module, name, None)
if original is None or getattr(original, '_graphiti_instrumented', False):
return
async def wrapped(*args, **kwargs):
request_id = main.CURRENT_REQUEST_ID.get() or 'unbound'
run_id = main.CURRENT_RUN_ID.get()
started = main._emit_event(request_id, run_id, phase, 'start')
phase_token = main.CURRENT_GRAPHITI_PHASE.set(phase)
try:
if phase == 'edge_resolution':
extracted_edges = args[1] if len(args) > 1 else kwargs.get('extracted_edges', [])
edge_total = len(extracted_edges or [])
stage_started = time.monotonic()
main._emit_event(request_id, run_id, phase, 'start', started,
timeout_seconds=main.GRAPHITI_EDGE_RESOLUTION_TIMEOUT_SECONDS,
details={'substage':'edge_resolution_enter','edge_index':0,'edge_total':edge_total,
'candidate_index':0,'candidate_total':None,
'elapsed_seconds':0.0,'last_progress_at':datetime.now(timezone.utc).isoformat()})
try:
result = await original(*args, **kwargs)
except asyncio.TimeoutError as exc:
elapsed = time.monotonic() - stage_started
main._emit_event(request_id, run_id, phase, 'timeout', started,
error='EdgeResolutionTimeout',
details={'substage':'edge_resolution_wait','edge_index':None,'edge_total':edge_total,
'candidate_index':None,'candidate_total':None,
'elapsed_seconds':round(elapsed,3),
'last_progress_at':datetime.now(timezone.utc).isoformat()})
raise main.EdgeResolutionTimeout(elapsed, edge_total) from exc
main._emit_event(request_id, run_id, phase, 'completed', started,
details={'substage':'edge_resolution_complete','edge_index':edge_total,'edge_total':edge_total,
'candidate_index':None,'candidate_total':None,
'elapsed_seconds':round(time.monotonic()-stage_started,3),
'last_progress_at':datetime.now(timezone.utc).isoformat()})
return result
result = await original(*args, **kwargs)
main._emit_event(request_id, run_id, phase, 'completed', started)
return result
except Exception as exc:
main._emit_event(request_id, run_id, phase, 'failed', started, error=exc)
raise
finally:
main.CURRENT_GRAPHITI_PHASE.reset(phase_token)
wrapped._graphiti_instrumented = True
setattr(graphiti_module, name, wrapped)
def _worker_phase_event(phase, started_at=None, outcome='start'):
request_id = main.CURRENT_REQUEST_ID.get() or 'unbound'
run_id = main.CURRENT_RUN_ID.get()
return main._emit_event(request_id, run_id, phase, outcome, started_at)
def install_sequential_graphiti_path():
"""Serialize only Graphiti's entity-resolution/edge-extraction branch."""
original_resolve_extracted_node = graphiti_module.resolve_extracted_nodes.__globals__['resolve_extracted_node']
edge_globals = graphiti_module.resolve_extracted_edges.__globals__
original_resolve_extracted_edge = edge_globals['resolve_extracted_edge']
edge_source = textwrap.dedent(inspect.getsource(original_resolve_extracted_edge))
edge_source = edge_source.replace('async def resolve_extracted_edge(',
'async def sequential_resolve_extracted_edge(', 1)
old_edge_gather = """resolved_edge, (valid_at, invalid_at), invalidation_candidates = await asyncio.gather(
dedupe_extracted_edge(llm_client, extracted_edge, related_edges),
extract_edge_dates(llm_client, extracted_edge, current_episode, previous_episodes),
get_edge_contradictions(llm_client, extracted_edge, existing_edges),
)"""
new_edge_sequential = """resolved_edge, (valid_at, invalid_at), invalidation_candidates = await asyncio.gather(
dedupe_extracted_edge(llm_client, extracted_edge, related_edges),
extract_edge_dates(llm_client, extracted_edge, current_episode, previous_episodes),
get_edge_contradictions(llm_client, extracted_edge, existing_edges),
)"""
assert old_edge_gather in edge_source, 'edge gather target missing'
edge_source = edge_source.replace(old_edge_gather, new_edge_sequential, 1)
exec(compile(edge_source, '<graphiti-sequential-edge-resolution>', 'exec'), edge_globals)
sequential_resolve_extracted_edge = edge_globals['sequential_resolve_extracted_edge']
if not getattr(original_resolve_extracted_edge, '_edge_progress_wrapped', False):
for function_name, operation in (
('dedupe_extracted_edge', 'dedupe'),
('extract_edge_dates', 'dates'),
('get_edge_contradictions', 'contradictions'),
):
original_operation = edge_globals[function_name]
async def operation_wrapper(*args, _original=original_operation, _operation=operation, **kwargs):
request_id = main.CURRENT_REQUEST_ID.get() or 'unbound'
run_id = main.CURRENT_RUN_ID.get()
edge_index = main.CURRENT_EDGE_INDEX.get()
edge_total = main.CURRENT_EDGE_TOTAL.get()
op_started = main._emit_event(request_id, run_id, 'edge_operation_start', 'start',
details={'edge_index':edge_index,'edge_total':edge_total,'operation':_operation,
'last_progress_at':datetime.now(timezone.utc).isoformat()})
op_token = main.CURRENT_EDGE_OPERATION.set(_operation)
try:
result = await _original(*args, **kwargs)
main._emit_event(request_id, run_id, 'edge_operation_completed', 'completed', op_started,
details={'edge_index':edge_index,'edge_total':edge_total,'operation':_operation,
'last_progress_at':datetime.now(timezone.utc).isoformat()})
return result
except Exception as exc:
main._emit_event(request_id, run_id, 'edge_operation_failed', 'failed', op_started,
error=exc, details={'edge_index':edge_index,'edge_total':edge_total,'operation':_operation,
'last_progress_at':datetime.now(timezone.utc).isoformat()})
raise
finally:
main.CURRENT_EDGE_OPERATION.reset(op_token)
edge_globals[function_name] = operation_wrapper
original_resolve_extracted_edge._edge_progress_wrapped = True
async def sequential_resolve_extracted_nodes(llm_client, extracted_nodes, existing_nodes_lists):
import asyncio as _aio
results = await _aio.gather(*[
original_resolve_extracted_node(llm_client, extracted_node, existing_nodes)
for extracted_node, existing_nodes in zip(extracted_nodes, existing_nodes_lists)
])
uuid_map = {}
resolved_nodes = []
for result in results:
uuid_map.update(result[1])
resolved_nodes.append(result[0])
return resolved_nodes, uuid_map
setattr(graphiti_module, 'resolve_extracted_nodes', sequential_resolve_extracted_nodes)
async def sequential_resolve_extracted_edges(llm_client, extracted_edges, related_edges_lists,
existing_edges_lists, current_episode, previous_episodes):
resolved_edges=[]; invalidated_edges=[]
edge_total=len(extracted_edges or [])
for index,(extracted_edge, related_edges, existing_edges) in enumerate(
zip(extracted_edges, related_edges_lists, existing_edges_lists), start=1):
request_id=main.CURRENT_REQUEST_ID.get() or 'unbound'; run_id=main.CURRENT_RUN_ID.get()
idx_token=main.CURRENT_EDGE_INDEX.set(index)
total_token=main.CURRENT_EDGE_TOTAL.set(edge_total)
started=main._emit_event(request_id,run_id,'edge_start','start',
details={'edge_index':index,'edge_total':edge_total,'last_progress_at':datetime.now(timezone.utc).isoformat()})
try:
result=await sequential_resolve_extracted_edge(llm_client, extracted_edge, related_edges,
existing_edges, current_episode, previous_episodes)
resolved_edges.append(result[0]); invalidated_edges.extend(result[1])
main._emit_event(request_id,run_id,'edge_completed','completed',started,
details={'edge_index':index,'edge_total':edge_total,'last_progress_at':datetime.now(timezone.utc).isoformat()})
except Exception as exc:
main._emit_event(request_id,run_id,'edge_operation_failed','failed',started,error=exc,
details={'edge_index':index,'edge_total':edge_total,'operation':'edge','last_progress_at':datetime.now(timezone.utc).isoformat()})
raise
finally:
main.CURRENT_EDGE_INDEX.reset(idx_token); main.CURRENT_EDGE_TOTAL.reset(total_token)
return resolved_edges, invalidated_edges
graphiti_module.resolve_extracted_edges = sequential_resolve_extracted_edges
source = textwrap.dedent(inspect.getsource(graphiti_module.Graphiti.add_episode))
source = source.replace('async def add_episode(', 'async def add_episode_sequential(', 1)
old = '''(mentioned_nodes, uuid_map), extracted_edges = await asyncio.gather(
resolve_extracted_nodes(self.llm_client, extracted_nodes, existing_nodes_lists),
extract_edges(
self.llm_client, episode, extracted_nodes, previous_episodes, group_id
),
)'''
new = '''resolution_started = _worker_phase_event("entity_resolution")
mentioned_nodes, uuid_map = await resolve_extracted_nodes(
self.llm_client, extracted_nodes, existing_nodes_lists
)
_worker_phase_event("entity_resolution", resolution_started, "completed")
extraction_started = _worker_phase_event("edge_extraction")
extracted_edges = await extract_edges(
self.llm_client, episode, extracted_nodes, previous_episodes, group_id
)
_worker_phase_event("edge_extraction", extraction_started, "completed")'''
if old not in source:
raise RuntimeError('sequential_graphiti_patch_target_not_found')
source = source.replace(old, new, 1)
namespace = graphiti_module.__dict__
namespace['_worker_phase_event'] = _worker_phase_event
exec(compile(source, '<graphiti-sequential-add-episode>', 'exec'), namespace)
graphiti_module.Graphiti.add_episode = namespace['add_episode_sequential']
def install_instrumentation():
install_sequential_graphiti_path()
for name, phase in (
('extract_nodes', 'entity_extraction'),
('resolve_extracted_nodes', 'entity_resolution'),
('extract_edges', 'edge_extraction'),
('resolve_extracted_edges', 'edge_resolution'),
('get_relevant_nodes', 'neo4j_read_nodes'),
('get_relevant_edges', 'neo4j_read_edges'),
('add_nodes_and_edges_bulk', 'neo4j_write_commit'),
):
instrument_async(name, phase)
async def run(payload):
request_id = payload['request_id']
run_id = payload.get('run_id', '')
main.CURRENT_REQUEST_ID.set(request_id)
main.CURRENT_RUN_ID.set(run_id)
main.CURRENT_WORKER_DEADLINE.set(payload.get('worker_deadline_monotonic'))
install_instrumentation()
client = main.build_graphiti_client()
original_retrieve = client.retrieve_episodes
async def retrieve_with_events(self, *args, **kwargs):
started = main._emit_event(request_id, run_id, 'neo4j_read_previous_episodes', 'start', timeout_seconds=5)
try:
result = await original_retrieve(*args, **kwargs)
main._emit_event(request_id, run_id, 'neo4j_read_previous_episodes', 'completed', started)
return result
except Exception as exc:
main._emit_event(request_id, run_id, 'neo4j_read_previous_episodes', 'failed', started, error=exc)
raise
client.retrieve_episodes = MethodType(retrieve_with_events, client)
started = main._emit_event(request_id, run_id, 'graphiti_processing', 'start', timeout_seconds=main.GRAPHITI_WORK_DEADLINE_SECONDS)
try:
result = await client.add_episode(
name=payload['name'],
episode_body=payload['content'],
source_description=payload['source_description'],
reference_time=__import__('datetime').datetime.fromisoformat(payload['timestamp']),
source=__import__('graphiti_core.nodes', fromlist=['EpisodeType']).EpisodeType.text,
)
main._emit_event(request_id, run_id, 'graphiti_processing', 'completed', started)
print('GRAPHITI_WORKER_RESULT=' + json.dumps({'ok': True, 'name': result.episode.name, 'uuid': str(result.episode.uuid), 'nodes': len(result.nodes), 'edges': len(result.edges)}), flush=True)
except Exception as exc:
main._emit_event(request_id, run_id, 'graphiti_processing', 'failed', started, error=exc)
print('GRAPHITI_WORKER_RESULT=' + json.dumps({'ok': False, 'error_type': type(exc).__name__,
'error_phase': main.LAST_GRAPHITI_PHASE.get() or main.CURRENT_GRAPHITI_PHASE.get() or 'graphiti_processing',
'timeout_reason': getattr(exc, 'timeout_reason', None),
'error': str(exc)[:500]}), flush=True)
raise
def main_entry():
payload = json.load(sys.stdin)
try:
asyncio.run(run(payload))
except Exception:
sys.exit(1)
if __name__ == '__main__':
main_entry()