Explorer
/tmp/worker_in_container.py
← Zurück ↓ Download
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()