diff --git a/lypi-agent-core/src/main/java/cn/lypi/agent/DefaultTurnExecutor.java b/lypi-agent-core/src/main/java/cn/lypi/agent/DefaultTurnExecutor.java index 8ff3b8df..8bc87f11 100644 --- a/lypi-agent-core/src/main/java/cn/lypi/agent/DefaultTurnExecutor.java +++ b/lypi-agent-core/src/main/java/cn/lypi/agent/DefaultTurnExecutor.java @@ -15,6 +15,10 @@ import cn.lypi.contracts.context.ToolCallContentBlock; import cn.lypi.contracts.event.ErrorEvent; import cn.lypi.contracts.event.TurnStartEvent; +import cn.lypi.contracts.hook.AfterTurnHookContext; +import cn.lypi.contracts.hook.BeforeTurnHookContext; +import cn.lypi.contracts.hook.BeforeTurnHookResult; +import cn.lypi.contracts.hook.TurnHookRuntime; import cn.lypi.contracts.model.AssistantEventStream; import cn.lypi.contracts.model.AssistantError; import cn.lypi.contracts.model.AssistantStart; @@ -44,8 +48,13 @@ public final class DefaultTurnExecutor implements TurnExecutor { private final ContextBudgetEstimator budgetEstimator; private final TurnContinuationGuard continuationGuard; private final TurnEventPublisher eventPublisher; + private final TurnHookRuntime turnHooks; public DefaultTurnExecutor(AgentCoreRuntimePorts ports, TurnIds ids, Clock clock) { + this(ports, ids, clock, TurnHookRuntime.noop()); + } + + public DefaultTurnExecutor(AgentCoreRuntimePorts ports, TurnIds ids, Clock clock, TurnHookRuntime turnHooks) { this.ports = ports; this.ids = ids; this.clock = clock; @@ -55,6 +64,7 @@ public DefaultTurnExecutor(AgentCoreRuntimePorts ports, TurnIds ids, Clock clock this.budgetEstimator = new ContextBudgetEstimator(); this.continuationGuard = new TurnContinuationGuard(ports.sessionManager()); this.eventPublisher = new TurnEventPublisher(ports.eventBus(), clock); + this.turnHooks = turnHooks == null ? TurnHookRuntime.noop() : turnHooks; } @Override @@ -73,30 +83,43 @@ private TurnState executeWithTurnId(TurnRequest request, String turnId) { request.parentEntryId().ifPresent(parentEntryId -> ports.sessionManager().switchLeaf(parentEntryId)); Instant startedAt = clock.instant(); ports.eventBus().publish(new TurnStartEvent(request.sessionId(), turnId, startedAt, startedAt)); - Optional unsafeReason = continuationGuard.unsafeContinuationReason(currentLeafId()); - if (unsafeReason.isPresent()) { - ports.eventBus().publish(new ErrorEvent( - request.sessionId(), - unsafeReason.orElseThrow(), - "当前分支停在 assistant 工具调用消息上,不能直接追加用户消息。请选择上一条用户消息或工具结果之后继续。", - clock.instant() - )); - return failedState(turnId, request.sessionId(), null, List.of(), 0, startedAt, currentLeafId()); - } - - AgentMessage user = messageFactory.userMessage(ids.newMessageId(), request.userInput()); - String contextLeafId = appendNewMessage(request.sessionId(), user); - newMessages.add(user); ContextSnapshot context = null; int toolRound = 0; + String contextLeafId = currentLeafId(); try { + BeforeTurnHookResult beforeHook = turnHooks.beforeTurn(new BeforeTurnHookContext(request, turnId, ports.cwd())); + if (beforeHook != null && beforeHook.blocked()) { + String message = beforeHook.message() == null ? "turn hook blocked" : beforeHook.message(); + ports.eventBus().publish(new ErrorEvent( + request.sessionId(), + "turn-hook-blocked", + message, + clock.instant() + )); + return failedState(request, turnId, null, List.of(), 0, startedAt, contextLeafId); + } + Optional unsafeReason = continuationGuard.unsafeContinuationReason(currentLeafId()); + if (unsafeReason.isPresent()) { + ports.eventBus().publish(new ErrorEvent( + request.sessionId(), + unsafeReason.orElseThrow(), + "当前分支停在 assistant 工具调用消息上,不能直接追加用户消息。请选择上一条用户消息或工具结果之后继续。", + clock.instant() + )); + return failedState(request, turnId, null, List.of(), 0, startedAt, currentLeafId()); + } + + AgentMessage user = messageFactory.userMessage(ids.newMessageId(), request.userInput()); + contextLeafId = appendNewMessage(request.sessionId(), user); + newMessages.add(user); + context = buildContext(request, Optional.of(contextLeafId)); AgentMessage assistant = runModel(request, context); contextLeafId = appendStartedMessage(request.sessionId(), assistant); newMessages.add(assistant); if (isAssistantError(assistant, request)) { - return failedState(turnId, request.sessionId(), context, newMessages, toolRound, startedAt, contextLeafId); + return failedState(request, turnId, context, newMessages, toolRound, startedAt, contextLeafId); } while (!request.abortSignal().aborted()) { @@ -106,9 +129,9 @@ private TurnState executeWithTurnId(TurnRequest request, String turnId) { "incomplete-tool-call", "模型返回的工具调用参数未完成,已终止本轮执行。" ); - appendNewMessage(request.sessionId(), error); + contextLeafId = appendNewMessage(request.sessionId(), error); newMessages.add(error); - return failedState(turnId, request.sessionId(), context, newMessages, toolRound, startedAt, contextLeafId); + return failedState(request, turnId, context, newMessages, toolRound, startedAt, contextLeafId); } List toolRequests = toolCallMapper.requestsFrom(assistant); if (toolRequests.isEmpty()) { @@ -134,7 +157,7 @@ private TurnState executeWithTurnId(TurnRequest request, String turnId) { contextLeafId = appendStartedMessage(request.sessionId(), assistant); newMessages.add(assistant); if (isAssistantError(assistant, request)) { - return failedState(turnId, request.sessionId(), context, newMessages, toolRound, startedAt, contextLeafId); + return failedState(request, turnId, context, newMessages, toolRound, startedAt, contextLeafId); } } } catch (RuntimeException failure) { @@ -145,13 +168,12 @@ private TurnState executeWithTurnId(TurnRequest request, String turnId) { ); appendNewMessage(request.sessionId(), handled.message()); newMessages.add(handled.message()); - return failedState(turnId, request.sessionId(), context, newMessages, toolRound, startedAt, currentLeafId()); + return failedState(request, turnId, context, newMessages, toolRound, startedAt, currentLeafId()); } TurnStatus status = request.abortSignal().aborted() ? TurnStatus.ABORTED : TurnStatus.COMPLETED; TurnState state = new TurnState(turnId, request.sessionId(), context, List.copyOf(newMessages), toolRound, status); - eventPublisher.publishTurnEnd(request.sessionId(), turnId, status, startedAt, toolRound, contextLeafId); - return state; + return finishTurn(request, state, startedAt, contextLeafId); } private boolean isAssistantError(AgentMessage assistant, TurnRequest request) { @@ -171,8 +193,8 @@ private String currentLeafId() { } private TurnState failedState( + TurnRequest request, String turnId, - String sessionId, ContextSnapshot context, List newMessages, int toolRound, @@ -181,14 +203,46 @@ private TurnState failedState( ) { TurnState state = new TurnState( turnId, - sessionId, + request.sessionId(), context, List.copyOf(newMessages), toolRound, TurnStatus.FAILED ); - eventPublisher.publishTurnEnd(sessionId, turnId, TurnStatus.FAILED, startedAt, toolRound, leafEntryId); - return state; + return finishTurn(request, state, startedAt, leafEntryId); + } + + private TurnState finishTurn(TurnRequest request, TurnState state, Instant startedAt, String leafEntryId) { + TurnState finalState = state; + try { + turnHooks.afterTurn(new AfterTurnHookContext(request, state, ports.cwd())); + } catch (RuntimeException failure) { + AgentCoreExceptionHandler.Failure handled = exceptionHandler.handle( + request.sessionId(), + ids.newMessageId(), + failure + ); + appendNewMessage(request.sessionId(), handled.message()); + List messages = new ArrayList<>(state.newMessages()); + messages.add(handled.message()); + finalState = new TurnState( + state.turnId(), + state.sessionId(), + state.context(), + List.copyOf(messages), + state.currentToolRound(), + TurnStatus.FAILED + ); + } + eventPublisher.publishTurnEnd( + request.sessionId(), + finalState.turnId(), + finalState.status(), + startedAt, + finalState.currentToolRound(), + leafEntryId + ); + return finalState; } private ContextSnapshot buildContext(TurnRequest request, Optional leafEntryId) { diff --git a/lypi-agent-core/src/test/java/cn/lypi/agent/DefaultTurnExecutorTest.java b/lypi-agent-core/src/test/java/cn/lypi/agent/DefaultTurnExecutorTest.java index 448b5902..90896b80 100644 --- a/lypi-agent-core/src/test/java/cn/lypi/agent/DefaultTurnExecutorTest.java +++ b/lypi-agent-core/src/test/java/cn/lypi/agent/DefaultTurnExecutorTest.java @@ -27,6 +27,10 @@ import cn.lypi.contracts.event.ToolEndEvent; import cn.lypi.contracts.event.ToolStartEvent; import cn.lypi.contracts.common.JsonSchema; +import cn.lypi.contracts.hook.AfterTurnHookResult; +import cn.lypi.contracts.hook.BeforeTurnHookResult; +import cn.lypi.contracts.hook.DefaultTurnHookRuntime; +import cn.lypi.contracts.hook.TurnHook; import cn.lypi.contracts.model.AssistantDone; import cn.lypi.contracts.model.AssistantError; import cn.lypi.contracts.model.AssistantStart; @@ -57,6 +61,196 @@ import static org.assertj.core.api.Assertions.assertThat; class DefaultTurnExecutorTest { + @Test + void beforeTurnHookCanBlockBeforeAppendingUserMessage() { + AgentCoreTestFixtures.InMemorySessionManager session = new AgentCoreTestFixtures.InMemorySessionManager(); + AgentCoreTestFixtures.StubAiProvider provider = new AgentCoreTestFixtures.StubAiProvider(); + AgentCoreTestFixtures.StubToolRuntime tools = new AgentCoreTestFixtures.StubToolRuntime(); + AgentCoreTestFixtures.RecordingEventBus eventBus = new AgentCoreTestFixtures.RecordingEventBus(); + DefaultTurnExecutor executor = new DefaultTurnExecutor( + AgentCoreTestFixtures.ports( + session, + provider, + tools, + eventBus, + request -> { + throw new AssertionError("阻断后不应构建 context"); + }, + new NoopCompactionCoordinator(), + new NoopMemoryExtractionWorker() + ), + TurnIds.fixed("turn-1", "msg-user", "entry-1"), + Clock.fixed(NOW, ZoneOffset.UTC), + new DefaultTurnHookRuntime(List.of(TurnHook.before(context -> BeforeTurnHookResult.block("denied")))) + ); + + TurnState state = executor.execute(new TurnRequest( + "session-1", + "hello", + Optional.empty(), + () -> false + )); + + assertThat(state.status()).isEqualTo(TurnStatus.FAILED); + assertThat(session.messages()).isEmpty(); + assertThat(provider.contexts).isEmpty(); + assertThat(eventBus.events).extracting(AgentEvent::getClass) + .containsExactly(TurnStartEvent.class, ErrorEvent.class, TurnEndEvent.class); + ErrorEvent error = (ErrorEvent) eventBus.events.get(1); + assertThat(error.message()).isEqualTo("denied"); + TurnEndEvent turnEnd = (TurnEndEvent) eventBus.events.getLast(); + assertThat(turnEnd.status()).isEqualTo("FAILED"); + assertThat(turnEnd.leafEntryId()).isEmpty(); + } + + @Test + void afterTurnHookReceivesCompletedStateBeforeTurnEndEvent() { + AgentCoreTestFixtures.InMemorySessionManager session = new AgentCoreTestFixtures.InMemorySessionManager(); + AgentCoreTestFixtures.StubAiProvider provider = new AgentCoreTestFixtures.StubAiProvider(); + AgentCoreTestFixtures.StubToolRuntime tools = new AgentCoreTestFixtures.StubToolRuntime(); + AgentCoreTestFixtures.RecordingEventBus eventBus = new AgentCoreTestFixtures.RecordingEventBus(); + Clock clock = Clock.fixed(NOW, ZoneOffset.UTC); + List observed = new ArrayList<>(); + provider.enqueue(List.of( + new AssistantStart("msg-assistant"), + new TextDelta("hi"), + new AssistantDone(Optional.empty(), Optional.of("end_turn")) + )); + ContextAssembler assembler = request -> new ContextAssembly( + AgentCoreTestFixtures.minimalContext(session.messages()), + AgentCoreTestFixtures.emptyResources(), + List.of(), + List.of(), + List.of(), + false + ); + DefaultTurnExecutor executor = new DefaultTurnExecutor( + AgentCoreTestFixtures.ports( + session, + provider, + tools, + eventBus, + assembler, + new NoopCompactionCoordinator(), + new NoopMemoryExtractionWorker() + ), + TurnIds.fixed("turn-1", "msg-user", "entry-1"), + clock, + new DefaultTurnHookRuntime(List.of(TurnHook.after(context -> { + observed.add(context.state()); + return AfterTurnHookResult.keep(); + }))) + ); + + TurnState state = executor.execute(new TurnRequest("session-1", "hello", Optional.empty(), () -> false)); + + assertThat(state.status()).isEqualTo(TurnStatus.COMPLETED); + assertThat(observed).hasSize(1); + assertThat(observed.getFirst().status()).isEqualTo(TurnStatus.COMPLETED); + assertThat(eventBus.events.getLast()).isInstanceOf(TurnEndEvent.class); + TurnEndEvent turnEnd = (TurnEndEvent) eventBus.events.getLast(); + assertThat(turnEnd.status()).isEqualTo("COMPLETED"); + assertThat(turnEnd.leafEntryId()).isEqualTo("entry-msg-assistant"); + } + + @Test + void afterTurnHookFailureKeepsOriginalLeafEntryId() { + AgentCoreTestFixtures.InMemorySessionManager session = new AgentCoreTestFixtures.InMemorySessionManager(); + AgentCoreTestFixtures.StubAiProvider provider = new AgentCoreTestFixtures.StubAiProvider(); + AgentCoreTestFixtures.StubToolRuntime tools = new AgentCoreTestFixtures.StubToolRuntime(); + AgentCoreTestFixtures.RecordingEventBus eventBus = new AgentCoreTestFixtures.RecordingEventBus(); + Clock clock = Clock.fixed(NOW, ZoneOffset.UTC); + provider.enqueue(List.of( + new AssistantStart("msg-assistant"), + new TextDelta("hi"), + new AssistantDone(Optional.empty(), Optional.of("end_turn")) + )); + ContextAssembler assembler = request -> new ContextAssembly( + AgentCoreTestFixtures.minimalContext(session.messages()), + AgentCoreTestFixtures.emptyResources(), + List.of(), + List.of(), + List.of(), + false + ); + DefaultTurnExecutor executor = new DefaultTurnExecutor( + AgentCoreTestFixtures.ports( + session, + provider, + tools, + eventBus, + assembler, + new NoopCompactionCoordinator(), + new NoopMemoryExtractionWorker() + ), + TurnIds.fixed("turn-1", "msg-user", "msg-fallback", "msg-hook-error"), + clock, + new DefaultTurnHookRuntime(List.of(TurnHook.after(context -> { + throw new IllegalStateException("after hook down"); + }))) + ); + + TurnState state = executor.execute(new TurnRequest("session-1", "hello", Optional.empty(), () -> false)); + + assertThat(state.status()).isEqualTo(TurnStatus.FAILED); + assertThat(state.currentToolRound()).isZero(); + assertThat(state.newMessages()).extracting(AgentMessage::id) + .containsExactly("msg-user", "msg-assistant", "msg-hook-error"); + assertThat(session.messages()).extracting(AgentMessage::id) + .containsExactly("msg-user", "msg-assistant", "msg-hook-error"); + TurnEndEvent turnEnd = (TurnEndEvent) eventBus.events.getLast(); + assertThat(turnEnd.status()).isEqualTo("FAILED"); + assertThat(turnEnd.toolRounds()).isZero(); + assertThat(turnEnd.leafEntryId()).isEqualTo("entry-msg-assistant"); + } + + @Test + void afterTurnHookReceivesFailedStateBeforeTurnEndEvent() { + AgentCoreTestFixtures.InMemorySessionManager session = new AgentCoreTestFixtures.InMemorySessionManager(); + AgentCoreTestFixtures.StubAiProvider provider = new AgentCoreTestFixtures.StubAiProvider(); + AgentCoreTestFixtures.StubToolRuntime tools = new AgentCoreTestFixtures.StubToolRuntime(); + AgentCoreTestFixtures.RecordingEventBus eventBus = new AgentCoreTestFixtures.RecordingEventBus(); + Clock clock = Clock.fixed(NOW, ZoneOffset.UTC); + List observed = new ArrayList<>(); + provider.failWith(new RuntimeException("provider down")); + ContextAssembler assembler = request -> new ContextAssembly( + AgentCoreTestFixtures.minimalContext(session.messages()), + AgentCoreTestFixtures.emptyResources(), + List.of(), + List.of(), + List.of(), + false + ); + DefaultTurnExecutor executor = new DefaultTurnExecutor( + AgentCoreTestFixtures.ports( + session, + provider, + tools, + eventBus, + assembler, + new NoopCompactionCoordinator(), + new NoopMemoryExtractionWorker() + ), + TurnIds.fixed("turn-1", "msg-user", "entry-1"), + clock, + new DefaultTurnHookRuntime(List.of(TurnHook.after(context -> { + observed.add(context.state()); + return AfterTurnHookResult.keep(); + }))) + ); + + TurnState state = executor.execute(new TurnRequest("session-1", "hello", Optional.empty(), () -> false)); + + assertThat(state.status()).isEqualTo(TurnStatus.FAILED); + assertThat(observed).hasSize(1); + assertThat(observed.getFirst().status()).isEqualTo(TurnStatus.FAILED); + assertThat(observed.getFirst().newMessages()).extracting(AgentMessage::id) + .containsExactly("msg-user", "entry-1"); + TurnEndEvent turnEnd = (TurnEndEvent) eventBus.events.getLast(); + assertThat(turnEnd.status()).isEqualTo("FAILED"); + assertThat(turnEnd.leafEntryId()).isEqualTo("entry-entry-1"); + } + @Test void executesSimpleTurnWithoutTools() { AgentCoreTestFixtures.InMemorySessionManager session = new AgentCoreTestFixtures.InMemorySessionManager(); @@ -1131,6 +1325,7 @@ void failsTurnWhenAssistantToolCallIsIncomplete() { "MessageStartEvent:msg-error", "MessageEndEvent:msg-error" ); + assertThat(((TurnEndEvent) eventBus.events.getLast()).leafEntryId()).isEqualTo("entry-msg-error"); assertThat(memory.calls).isZero(); } diff --git a/lypi-boot/src/main/java/cn/lypi/boot/runtime/LyPiRuntimeAutoConfiguration.java b/lypi-boot/src/main/java/cn/lypi/boot/runtime/LyPiRuntimeAutoConfiguration.java index 67f4e41c..2938dc3b 100644 --- a/lypi-boot/src/main/java/cn/lypi/boot/runtime/LyPiRuntimeAutoConfiguration.java +++ b/lypi-boot/src/main/java/cn/lypi/boot/runtime/LyPiRuntimeAutoConfiguration.java @@ -7,6 +7,7 @@ import cn.lypi.boot.tool.LyPiPermissionsProperties; import cn.lypi.boot.tool.ToolRuntimeFactoryPort; import cn.lypi.contracts.event.EventBus; +import cn.lypi.contracts.hook.TurnHook; import cn.lypi.contracts.model.ModelCatalogPort; import cn.lypi.contracts.runtime.AgentCenterPort; import cn.lypi.contracts.runtime.AgentCoreFactoryPort; @@ -222,7 +223,8 @@ public AgentCorePort agentCore( ResourceRuntimePort resourceRuntime, EventBus eventBus, ContextAssembler contextAssembler, - CompactionCoordinator compactionCoordinator + CompactionCoordinator compactionCoordinator, + ObjectProvider turnHooks ) { return RuntimeBeanFactories.agentCore( properties, @@ -234,6 +236,7 @@ public AgentCorePort agentCore( eventBus, contextAssembler, compactionCoordinator, + turnHooks, Clock.systemUTC() ); } @@ -255,6 +258,7 @@ public AgentCoreFactoryPort agentCoreFactory( EventBus eventBus, ObjectProvider compactionSummarizer, ObjectProvider modelCatalog, + ObjectProvider turnHooks, Clock clock ) { return RuntimeBeanFactories.agentCoreFactory( @@ -266,6 +270,7 @@ public AgentCoreFactoryPort agentCoreFactory( eventBus, compactionSummarizer, modelCatalog, + turnHooks, clock ); } diff --git a/lypi-boot/src/main/java/cn/lypi/boot/runtime/RuntimeBeanFactories.java b/lypi-boot/src/main/java/cn/lypi/boot/runtime/RuntimeBeanFactories.java index 0fb86599..d2899919 100644 --- a/lypi-boot/src/main/java/cn/lypi/boot/runtime/RuntimeBeanFactories.java +++ b/lypi-boot/src/main/java/cn/lypi/boot/runtime/RuntimeBeanFactories.java @@ -20,6 +20,8 @@ import cn.lypi.boot.tool.ToolRuntimeFactoryPort; import cn.lypi.contracts.context.ContextBudget; import cn.lypi.contracts.event.EventBus; +import cn.lypi.contracts.hook.DefaultTurnHookRuntime; +import cn.lypi.contracts.hook.TurnHook; import cn.lypi.contracts.model.ModelCatalogPort; import cn.lypi.contracts.model.ModelSelection; import cn.lypi.contracts.runtime.AgentCenterPort; @@ -216,6 +218,7 @@ static AgentCorePort agentCore( EventBus eventBus, ContextAssembler contextAssembler, CompactionCoordinator compactionCoordinator, + ObjectProvider turnHooks, Clock clock ) { AgentCoreRuntimePorts ports = new AgentCoreRuntimePorts( @@ -231,7 +234,7 @@ static AgentCorePort agentCore( compactionCoordinator, new NoopMemoryExtractionWorker() ); - return new DefaultTurnExecutor(ports, TurnIds.random(), clock); + return new DefaultTurnExecutor(ports, TurnIds.random(), clock, turnHookRuntime(turnHooks)); } static AgentCoreFactoryPort agentCoreFactory( @@ -243,6 +246,7 @@ static AgentCoreFactoryPort agentCoreFactory( EventBus eventBus, ObjectProvider compactionSummarizer, ObjectProvider modelCatalog, + ObjectProvider turnHooks, Clock clock ) { return new AgentCoreFactoryPort() { @@ -308,12 +312,17 @@ private AgentCorePort createWithPorts( new NoopMemoryExtractionWorker() ), TurnIds.random(), - clock + clock, + turnHookRuntime(turnHooks) ); } }; } + private static DefaultTurnHookRuntime turnHookRuntime(ObjectProvider turnHooks) { + return new DefaultTurnHookRuntime(turnHooks == null ? List.of() : turnHooks.orderedStream().toList()); + } + static MemoryConsolidationTrigger memoryConsolidationTrigger() { return new MemoryConsolidationTrigger(); } diff --git a/lypi-boot/src/main/java/cn/lypi/boot/tool/LyPiToolAutoConfiguration.java b/lypi-boot/src/main/java/cn/lypi/boot/tool/LyPiToolAutoConfiguration.java index 1cde4632..9b87f87d 100644 --- a/lypi-boot/src/main/java/cn/lypi/boot/tool/LyPiToolAutoConfiguration.java +++ b/lypi-boot/src/main/java/cn/lypi/boot/tool/LyPiToolAutoConfiguration.java @@ -3,6 +3,8 @@ import cn.lypi.contracts.runtime.AgentCenterPort; import cn.lypi.contracts.runtime.AgentRegistryPort; import cn.lypi.contracts.event.EventBus; +import cn.lypi.contracts.hook.DefaultToolHookRuntime; +import cn.lypi.contracts.hook.ToolHook; import cn.lypi.contracts.mcp.McpTransport; import cn.lypi.contracts.runtime.Executor; import cn.lypi.contracts.runtime.MailboxPort; @@ -27,6 +29,7 @@ import cn.lypi.tool.PermissionGate; import cn.lypi.tool.PermissionPromptPort; import cn.lypi.tool.PermissionResponseGate; +import cn.lypi.tool.ToolHookExecutionInterceptor; import cn.lypi.tool.ToolRuntimeOptions; import cn.lypi.tool.builtin.BuiltInTools; import cn.lypi.tool.mcp.McpClientManager; @@ -162,6 +165,7 @@ public ToolRuntimeFactoryPort toolRuntimeFactory( ObjectProvider agentCenter, ObjectProvider mailbox, ObjectProvider agentRegistry, + ObjectProvider toolHooks, SandboxPolicyResolver sandboxPolicyResolver, ObjectProvider eventBus, ObjectProvider responseGate, @@ -218,7 +222,7 @@ private ToolRuntimePort createRuntime( new cn.lypi.tool.ToolExecutionPlanner(), new cn.lypi.tool.ToolResultBudgeter(), new cn.lypi.tool.ToolRuntimeContextFactory(options), - cn.lypi.tool.ToolExecutionInterceptors.noop(), + new ToolHookExecutionInterceptor(new DefaultToolHookRuntime(toolHooks.orderedStream().toList(), runtimeEventBus)), securityRuntime, runtimeResponseGate, runtimePromptPort, diff --git a/lypi-boot/src/test/java/cn/lypi/boot/runtime/LyPiRuntimeAutoConfigurationTest.java b/lypi-boot/src/test/java/cn/lypi/boot/runtime/LyPiRuntimeAutoConfigurationTest.java index 1b0bfbd4..c2b79f31 100644 --- a/lypi-boot/src/test/java/cn/lypi/boot/runtime/LyPiRuntimeAutoConfigurationTest.java +++ b/lypi-boot/src/test/java/cn/lypi/boot/runtime/LyPiRuntimeAutoConfigurationTest.java @@ -28,6 +28,8 @@ import cn.lypi.contracts.event.PermissionRequestEvent; import cn.lypi.contracts.event.PermissionResponseEvent; import cn.lypi.contracts.event.SessionStateEvent; +import cn.lypi.contracts.hook.BeforeTurnHookResult; +import cn.lypi.contracts.hook.TurnHook; import cn.lypi.contracts.model.AssistantDone; import cn.lypi.contracts.model.AssistantEventStream; import cn.lypi.contracts.model.AssistantStart; @@ -136,6 +138,7 @@ import java.util.concurrent.CountDownLatch; import java.util.concurrent.TimeUnit; import java.util.concurrent.atomic.AtomicReference; +import java.util.concurrent.atomic.AtomicBoolean; import org.junit.jupiter.api.Test; import org.junit.jupiter.api.io.TempDir; import org.springframework.boot.ApplicationRunner; @@ -1277,6 +1280,68 @@ void createsAgentCoreFactoryWhenRequiredRuntimePortsExist() { .run(context -> assertThat(context).hasSingleBean(AgentCoreFactoryPort.class)); } + @Test + void agentCoreFactoryWiresTurnHooks() { + AtomicBoolean hookCalled = new AtomicBoolean(false); + + new ApplicationContextRunner() + .withUserConfiguration(LyPiRuntimeAutoConfiguration.class) + .withBean(AiProviderRuntimePort.class, () -> (snapshot, signal) -> { + throw new AssertionError("阻断后不应调用模型"); + }) + .withBean(ToolRuntimePort.class, NoopToolRuntime::new) + .withBean(SecurityRuntimePort.class, () -> LyPiRuntimeAutoConfigurationTest::allowAllSecurity) + .withBean(ResourceRuntimePort.class, NoopResourceRuntime::new) + .withBean(CompactionSummarizer.class, () -> request -> new CompactSummaryResult( + "summary", + new TokenUsage(0, 0, 0, 0) + )) + .withBean(TurnHook.class, () -> TurnHook.before(context -> { + hookCalled.set(true); + return BeforeTurnHookResult.block("boot hook denied"); + })) + .run(context -> { + AgentCorePort core = context.getBean(AgentCoreFactoryPort.class) + .create(tempDir, context.getBean(SessionManagerPort.class)); + + TurnState state = core.execute(new TurnRequest("session-hook", "hello", Optional.empty(), () -> false)); + + assertThat(hookCalled).isTrue(); + assertThat(state.status()).isEqualTo(TurnStatus.FAILED); + assertThat(state.newMessages()).isEmpty(); + }); + } + + @Test + void defaultAgentCoreWiresTurnHooks() { + AtomicBoolean hookCalled = new AtomicBoolean(false); + + new ApplicationContextRunner() + .withUserConfiguration(LyPiRuntimeAutoConfiguration.class) + .withBean(AiProviderRuntimePort.class, () -> (snapshot, signal) -> { + throw new AssertionError("阻断后不应调用模型"); + }) + .withBean(ToolRuntimePort.class, NoopToolRuntime::new) + .withBean(SecurityRuntimePort.class, () -> LyPiRuntimeAutoConfigurationTest::allowAllSecurity) + .withBean(ResourceRuntimePort.class, NoopResourceRuntime::new) + .withBean(CompactionSummarizer.class, () -> request -> new CompactSummaryResult( + "summary", + new TokenUsage(0, 0, 0, 0) + )) + .withBean(TurnHook.class, () -> TurnHook.before(context -> { + hookCalled.set(true); + return BeforeTurnHookResult.block("boot hook denied"); + })) + .run(context -> { + TurnState state = context.getBean(AgentCorePort.class) + .execute(new TurnRequest("session-hook", "hello", Optional.empty(), () -> false)); + + assertThat(hookCalled).isTrue(); + assertThat(state.status()).isEqualTo(TurnStatus.FAILED); + assertThat(state.newMessages()).isEmpty(); + }); + } + @Test void createsBackgroundMemoryConsolidationBeansWhenRuntimePortsExist() { new ApplicationContextRunner() diff --git a/lypi-boot/src/test/java/cn/lypi/boot/tool/LyPiToolAutoConfigurationTest.java b/lypi-boot/src/test/java/cn/lypi/boot/tool/LyPiToolAutoConfigurationTest.java index b2ad6136..b862f3cd 100644 --- a/lypi-boot/src/test/java/cn/lypi/boot/tool/LyPiToolAutoConfigurationTest.java +++ b/lypi-boot/src/test/java/cn/lypi/boot/tool/LyPiToolAutoConfigurationTest.java @@ -13,12 +13,18 @@ import cn.lypi.contracts.event.EventConsumer; import cn.lypi.contracts.event.EventFilter; import cn.lypi.contracts.event.EventSubscription; +import cn.lypi.contracts.event.HookEndEvent; +import cn.lypi.contracts.event.HookStartEvent; import cn.lypi.contracts.event.PermissionDecisionEvent; import cn.lypi.contracts.event.PermissionRequestEvent; import cn.lypi.contracts.event.PermissionResponseEvent; import cn.lypi.contracts.event.ToolEndEvent; import cn.lypi.contracts.event.ToolProgressEvent; import cn.lypi.contracts.event.ToolStartEvent; +import cn.lypi.contracts.hook.AfterToolHookContext; +import cn.lypi.contracts.hook.AfterToolHookResult; +import cn.lypi.contracts.hook.BeforeToolHookResult; +import cn.lypi.contracts.hook.ToolHook; import cn.lypi.contracts.model.ModelSelection; import cn.lypi.contracts.model.ThinkingLevel; import cn.lypi.contracts.mcp.McpServerConfig; @@ -81,9 +87,12 @@ import java.util.concurrent.CompletableFuture; import java.util.concurrent.CountDownLatch; import java.util.concurrent.TimeUnit; +import java.util.concurrent.atomic.AtomicInteger; import java.util.concurrent.atomic.AtomicReference; +import org.springframework.beans.factory.config.ConfigurableBeanFactory; import org.junit.jupiter.api.Test; import org.springframework.boot.test.context.runner.ApplicationContextRunner; +import org.springframework.core.Ordered; import static org.assertj.core.api.Assertions.assertThat; @@ -466,6 +475,150 @@ void defaultRuntimePublishesToolLifecycleAndProgressWhenEventBusIsAvailable() { }); } + @Test + void toolRuntimeFactoryAppliesRegisteredToolHooks() { + new ApplicationContextRunner() + .withUserConfiguration(LyPiToolAutoConfiguration.class) + .withBean(SecurityRuntimePort.class, () -> LyPiToolAutoConfigurationTest::allowAllSecurity) + .withBean(ToolHook.class, () -> ToolHook.after(context -> AfterToolHookResult.replace( + new ToolResult<>("hooked", false, context.result().newMessages(), context.result().replacement()) + ))) + .run(context -> { + ToolRuntimeFactoryPort factory = context.getBean(ToolRuntimeFactoryPort.class); + ToolRuntimePort runtime = factory.create(Path.of(".")); + runtime.register(new ProgressTool()); + + ToolResult result = runtime.execute( + List.of(new ToolUseRequest("toolu_1", "progress-test", Map.of("text", "done"), "msg_1")), + context() + ).getFirst(); + + assertThat(result.output()).isEqualTo("hooked"); + }); + } + + @Test + void toolRuntimeFactoryPublishesHookLifecycleEvents() { + RecordingEventBus eventBus = new RecordingEventBus(); + + new ApplicationContextRunner() + .withUserConfiguration(LyPiToolAutoConfiguration.class) + .withBean(EventBus.class, () -> eventBus) + .withBean(SecurityRuntimePort.class, () -> LyPiToolAutoConfigurationTest::allowAllSecurity) + .withBean(ToolHook.class, () -> ToolHook.before(context -> BeforeToolHookResult.allow())) + .run(context -> { + ToolRuntimeFactoryPort factory = context.getBean(ToolRuntimeFactoryPort.class); + ToolRuntimePort runtime = factory.create(Path.of(".")); + runtime.register(new ProgressTool()); + + ToolResult result = runtime.execute( + List.of(new ToolUseRequest("toolu_1", "progress-test", Map.of("text", "done"), "msg_1")), + context() + ).getFirst(); + + assertThat(result.isError()).isFalse(); + ToolStartEvent toolStart = eventBus.events.stream() + .filter(ToolStartEvent.class::isInstance) + .map(ToolStartEvent.class::cast) + .findFirst() + .orElseThrow(); + HookStartEvent hookStart = eventBus.events.stream() + .filter(HookStartEvent.class::isInstance) + .map(HookStartEvent.class::cast) + .findFirst() + .orElseThrow(); + HookEndEvent hookEnd = eventBus.events.stream() + .filter(HookEndEvent.class::isInstance) + .map(HookEndEvent.class::cast) + .findFirst() + .orElseThrow(); + ToolEndEvent toolEnd = eventBus.events.stream() + .filter(ToolEndEvent.class::isInstance) + .map(ToolEndEvent.class::cast) + .findFirst() + .orElseThrow(); + + assertThat(hookStart.sessionId()).isEqualTo(toolStart.sessionId()); + assertThat(hookStart.toolUseId()).isEqualTo(toolStart.toolUseId()); + assertThat(hookStart.parentMessageId()).isEqualTo(toolStart.parentMessageId()); + assertThat(hookStart.turnId()).isEqualTo(toolStart.turnId()); + assertThat(hookStart.toolName()).isEqualTo(toolStart.toolName()); + assertThat(hookEnd.sessionId()).isEqualTo(toolEnd.sessionId()); + assertThat(hookEnd.toolUseId()).isEqualTo(toolEnd.toolUseId()); + assertThat(hookEnd.hookRunId()).isEqualTo(hookStart.hookRunId()); + }); + } + + @Test + void toolRuntimeFactoryKeepsOriginalResultWhenNoToolHooksAreRegistered() { + new ApplicationContextRunner() + .withUserConfiguration(LyPiToolAutoConfiguration.class) + .withBean(SecurityRuntimePort.class, () -> LyPiToolAutoConfigurationTest::allowAllSecurity) + .run(context -> { + ToolRuntimeFactoryPort factory = context.getBean(ToolRuntimeFactoryPort.class); + ToolRuntimePort runtime = factory.create(Path.of(".")); + runtime.register(new ProgressTool()); + + ToolResult result = runtime.execute( + List.of(new ToolUseRequest("toolu_1", "progress-test", Map.of("text", "done"), "msg_1")), + context() + ).getFirst(); + + assertThat(result.output()).isEqualTo("done"); + }); + } + + @Test + void toolRuntimeFactoryAppliesToolHooksInSpringOrder() { + new ApplicationContextRunner() + .withUserConfiguration(LyPiToolAutoConfiguration.class) + .withBean(SecurityRuntimePort.class, () -> LyPiToolAutoConfigurationTest::allowAllSecurity) + .withBean("lateHook", ToolHook.class, () -> new OrderedAppendHook(20, "-late")) + .withBean("earlyHook", ToolHook.class, () -> new OrderedAppendHook(10, "-early")) + .run(context -> { + ToolRuntimeFactoryPort factory = context.getBean(ToolRuntimeFactoryPort.class); + ToolRuntimePort runtime = factory.create(Path.of(".")); + runtime.register(new ProgressTool()); + + ToolResult result = runtime.execute( + List.of(new ToolUseRequest("toolu_1", "progress-test", Map.of("text", "done"), "msg_1")), + context() + ).getFirst(); + + assertThat(result.output()).isEqualTo("done-early-late"); + }); + } + + @Test + void toolRuntimeFactoryResolvesPrototypeToolHooksForEachRuntimeCreation() { + AtomicInteger hookCreations = new AtomicInteger(); + + new ApplicationContextRunner() + .withUserConfiguration(LyPiToolAutoConfiguration.class) + .withBean(SecurityRuntimePort.class, () -> LyPiToolAutoConfigurationTest::allowAllSecurity) + .withBean( + "prototypeHook", + ToolHook.class, + () -> { + hookCreations.incrementAndGet(); + return ToolHook.after(context -> AfterToolHookResult.keep()); + }, + beanDefinition -> beanDefinition.setScope(ConfigurableBeanFactory.SCOPE_PROTOTYPE) + ) + .run(context -> { + ToolRuntimeFactoryPort factory = context.getBean(ToolRuntimeFactoryPort.class); + int baselineCreations = hookCreations.get(); + + factory.create(Path.of("build/runtime-a")); + int afterFirstRuntime = hookCreations.get(); + factory.create(Path.of("build/runtime-b")); + int afterSecondRuntime = hookCreations.get(); + + assertThat(afterFirstRuntime).isGreaterThan(baselineCreations); + assertThat(afterSecondRuntime).isGreaterThan(afterFirstRuntime); + }); + } + @Test void registersSubagentToolsWhenRuntimePortsAreAvailable() { new ApplicationContextRunner() @@ -806,6 +959,29 @@ public void close() { } } + private static final class OrderedAppendHook implements ToolHook, Ordered { + private final int order; + private final String suffix; + + private OrderedAppendHook(int order, String suffix) { + this.order = order; + this.suffix = suffix; + } + + @Override + public AfterToolHookResult afterToolCall(AfterToolHookContext context) { + String output = String.valueOf(context.result().output()) + suffix; + return AfterToolHookResult.replace( + new ToolResult<>(output, false, context.result().newMessages(), context.result().replacement()) + ); + } + + @Override + public int getOrder() { + return order; + } + } + private static final class AskTool implements Tool, String> { @Override public String name() { diff --git a/lypi-contracts/src/main/java/cn/lypi/contracts/event/AgentEvent.java b/lypi-contracts/src/main/java/cn/lypi/contracts/event/AgentEvent.java index caa8ffe4..a7a5f390 100644 --- a/lypi-contracts/src/main/java/cn/lypi/contracts/event/AgentEvent.java +++ b/lypi-contracts/src/main/java/cn/lypi/contracts/event/AgentEvent.java @@ -15,6 +15,8 @@ @JsonSubTypes.Type(value = ToolStartEvent.class, name = "tool_start"), @JsonSubTypes.Type(value = ToolProgressEvent.class, name = "tool_progress"), @JsonSubTypes.Type(value = ToolEndEvent.class, name = "tool_end"), + @JsonSubTypes.Type(value = HookStartEvent.class, name = "hook_start"), + @JsonSubTypes.Type(value = HookEndEvent.class, name = "hook_end"), @JsonSubTypes.Type(value = PermissionRequestEvent.class, name = "permission_request"), @JsonSubTypes.Type(value = PermissionResponseEvent.class, name = "permission_response"), @JsonSubTypes.Type(value = PermissionDecisionEvent.class, name = "permission_decision"), @@ -37,6 +39,8 @@ public sealed interface AgentEvent permits ToolStartEvent, ToolProgressEvent, ToolEndEvent, + HookStartEvent, + HookEndEvent, PermissionRequestEvent, PermissionResponseEvent, PermissionDecisionEvent, diff --git a/lypi-contracts/src/main/java/cn/lypi/contracts/event/HookEndEvent.java b/lypi-contracts/src/main/java/cn/lypi/contracts/event/HookEndEvent.java new file mode 100644 index 00000000..42368c1e --- /dev/null +++ b/lypi-contracts/src/main/java/cn/lypi/contracts/event/HookEndEvent.java @@ -0,0 +1,37 @@ +package cn.lypi.contracts.event; + +import cn.lypi.contracts.hook.HookPhase; +import cn.lypi.contracts.hook.HookRunStatus; +import java.time.Instant; +import java.util.Objects; + +public record HookEndEvent( + String sessionId, + String toolUseId, + String parentMessageId, + String turnId, + String toolName, + String hookRunId, + String hookName, + HookPhase phase, + HookRunStatus status, + String message, + Instant startedAt, + Instant endedAt, + long durationMillis, + Instant timestamp +) implements AgentEvent { + public HookEndEvent { + sessionId = Objects.requireNonNull(sessionId, "sessionId"); + toolUseId = Objects.requireNonNull(toolUseId, "toolUseId"); + toolName = Objects.requireNonNull(toolName, "toolName"); + hookRunId = Objects.requireNonNull(hookRunId, "hookRunId"); + hookName = Objects.requireNonNull(hookName, "hookName"); + phase = Objects.requireNonNull(phase, "phase"); + status = Objects.requireNonNull(status, "status"); + startedAt = Objects.requireNonNull(startedAt, "startedAt"); + endedAt = Objects.requireNonNull(endedAt, "endedAt"); + durationMillis = Math.max(0L, durationMillis); + timestamp = Objects.requireNonNull(timestamp, "timestamp"); + } +} diff --git a/lypi-contracts/src/main/java/cn/lypi/contracts/event/HookStartEvent.java b/lypi-contracts/src/main/java/cn/lypi/contracts/event/HookStartEvent.java new file mode 100644 index 00000000..70eef92f --- /dev/null +++ b/lypi-contracts/src/main/java/cn/lypi/contracts/event/HookStartEvent.java @@ -0,0 +1,29 @@ +package cn.lypi.contracts.event; + +import cn.lypi.contracts.hook.HookPhase; +import java.time.Instant; +import java.util.Objects; + +public record HookStartEvent( + String sessionId, + String toolUseId, + String parentMessageId, + String turnId, + String toolName, + String hookRunId, + String hookName, + HookPhase phase, + Instant startedAt, + Instant timestamp +) implements AgentEvent { + public HookStartEvent { + sessionId = Objects.requireNonNull(sessionId, "sessionId"); + toolUseId = Objects.requireNonNull(toolUseId, "toolUseId"); + toolName = Objects.requireNonNull(toolName, "toolName"); + hookRunId = Objects.requireNonNull(hookRunId, "hookRunId"); + hookName = Objects.requireNonNull(hookName, "hookName"); + phase = Objects.requireNonNull(phase, "phase"); + startedAt = Objects.requireNonNull(startedAt, "startedAt"); + timestamp = Objects.requireNonNull(timestamp, "timestamp"); + } +} diff --git a/lypi-contracts/src/main/java/cn/lypi/contracts/hook/AfterToolHookContext.java b/lypi-contracts/src/main/java/cn/lypi/contracts/hook/AfterToolHookContext.java new file mode 100644 index 00000000..19b04aae --- /dev/null +++ b/lypi-contracts/src/main/java/cn/lypi/contracts/hook/AfterToolHookContext.java @@ -0,0 +1,33 @@ +package cn.lypi.contracts.hook; + +import cn.lypi.contracts.tool.Tool; +import cn.lypi.contracts.tool.ToolResult; +import cn.lypi.contracts.tool.ToolUseContext; +import cn.lypi.contracts.tool.ToolUseRequest; +import java.util.Map; +import java.util.Objects; + +/** + * 表示工具执行后 hook 的上下文。 + */ +public record AfterToolHookContext( + ToolUseRequest request, + Tool tool, + Map input, + ToolUseContext toolContext, + ToolResult result +) { + public AfterToolHookContext { + request = Objects.requireNonNull(request, "request"); + tool = Objects.requireNonNull(tool, "tool"); + input = ToolHookInputSnapshots.snapshot(Objects.requireNonNull(input, "input")); + request = new ToolUseRequest( + request.toolUseId(), + request.toolName(), + input, + request.parentMessageId() + ); + toolContext = Objects.requireNonNull(toolContext, "toolContext"); + result = Objects.requireNonNull(result, "result"); + } +} diff --git a/lypi-contracts/src/main/java/cn/lypi/contracts/hook/AfterToolHookResult.java b/lypi-contracts/src/main/java/cn/lypi/contracts/hook/AfterToolHookResult.java new file mode 100644 index 00000000..0f491ec1 --- /dev/null +++ b/lypi-contracts/src/main/java/cn/lypi/contracts/hook/AfterToolHookResult.java @@ -0,0 +1,30 @@ +package cn.lypi.contracts.hook; + +import cn.lypi.contracts.tool.ToolResult; +import java.util.Objects; +import java.util.Optional; + +/** + * 表示工具执行后 hook 的处理结果。 + */ +public record AfterToolHookResult( + Optional> replacement +) { + public AfterToolHookResult { + replacement = replacement == null ? Optional.empty() : replacement; + } + + /** + * 返回保留原始结果的处理结果。 + */ + public static AfterToolHookResult keep() { + return new AfterToolHookResult(Optional.empty()); + } + + /** + * 返回替换工具结果的处理结果。 + */ + public static AfterToolHookResult replace(ToolResult result) { + return new AfterToolHookResult(Optional.of(Objects.requireNonNull(result, "result"))); + } +} diff --git a/lypi-contracts/src/main/java/cn/lypi/contracts/hook/AfterTurnHookContext.java b/lypi-contracts/src/main/java/cn/lypi/contracts/hook/AfterTurnHookContext.java new file mode 100644 index 00000000..9a065c5a --- /dev/null +++ b/lypi-contracts/src/main/java/cn/lypi/contracts/hook/AfterTurnHookContext.java @@ -0,0 +1,21 @@ +package cn.lypi.contracts.hook; + +import cn.lypi.contracts.agent.TurnRequest; +import cn.lypi.contracts.agent.TurnState; +import java.nio.file.Path; +import java.util.Objects; + +/** + * 表示 turn 执行后 hook 的上下文。 + */ +public record AfterTurnHookContext( + TurnRequest request, + TurnState state, + Path cwd +) { + public AfterTurnHookContext { + request = Objects.requireNonNull(request, "request"); + state = Objects.requireNonNull(state, "state"); + cwd = Objects.requireNonNull(cwd, "cwd").toAbsolutePath().normalize(); + } +} diff --git a/lypi-contracts/src/main/java/cn/lypi/contracts/hook/AfterTurnHookResult.java b/lypi-contracts/src/main/java/cn/lypi/contracts/hook/AfterTurnHookResult.java new file mode 100644 index 00000000..01374ab9 --- /dev/null +++ b/lypi-contracts/src/main/java/cn/lypi/contracts/hook/AfterTurnHookResult.java @@ -0,0 +1,13 @@ +package cn.lypi.contracts.hook; + +/** + * 表示 turn 执行后 hook 的处理结果。 + */ +public record AfterTurnHookResult() { + /** + * 返回保留原始 turn 状态的处理结果。 + */ + public static AfterTurnHookResult keep() { + return new AfterTurnHookResult(); + } +} diff --git a/lypi-contracts/src/main/java/cn/lypi/contracts/hook/BeforeToolHookContext.java b/lypi-contracts/src/main/java/cn/lypi/contracts/hook/BeforeToolHookContext.java new file mode 100644 index 00000000..c8c383cc --- /dev/null +++ b/lypi-contracts/src/main/java/cn/lypi/contracts/hook/BeforeToolHookContext.java @@ -0,0 +1,30 @@ +package cn.lypi.contracts.hook; + +import cn.lypi.contracts.tool.Tool; +import cn.lypi.contracts.tool.ToolUseContext; +import cn.lypi.contracts.tool.ToolUseRequest; +import java.util.Map; +import java.util.Objects; + +/** + * 表示工具执行前 hook 的上下文。 + */ +public record BeforeToolHookContext( + ToolUseRequest request, + Tool tool, + Map input, + ToolUseContext toolContext +) { + public BeforeToolHookContext { + request = Objects.requireNonNull(request, "request"); + tool = Objects.requireNonNull(tool, "tool"); + input = ToolHookInputSnapshots.snapshot(Objects.requireNonNull(input, "input")); + request = new ToolUseRequest( + request.toolUseId(), + request.toolName(), + input, + request.parentMessageId() + ); + toolContext = Objects.requireNonNull(toolContext, "toolContext"); + } +} diff --git a/lypi-contracts/src/main/java/cn/lypi/contracts/hook/BeforeToolHookResult.java b/lypi-contracts/src/main/java/cn/lypi/contracts/hook/BeforeToolHookResult.java new file mode 100644 index 00000000..4128f88a --- /dev/null +++ b/lypi-contracts/src/main/java/cn/lypi/contracts/hook/BeforeToolHookResult.java @@ -0,0 +1,23 @@ +package cn.lypi.contracts.hook; + +/** + * 表示工具执行前 hook 的决策结果。 + */ +public record BeforeToolHookResult( + boolean blocked, + String message +) { + /** + * 返回允许继续执行的结果。 + */ + public static BeforeToolHookResult allow() { + return new BeforeToolHookResult(false, null); + } + + /** + * 返回阻断工具执行的结果。 + */ + public static BeforeToolHookResult block(String message) { + return new BeforeToolHookResult(true, message); + } +} diff --git a/lypi-contracts/src/main/java/cn/lypi/contracts/hook/BeforeTurnHookContext.java b/lypi-contracts/src/main/java/cn/lypi/contracts/hook/BeforeTurnHookContext.java new file mode 100644 index 00000000..2147dfd5 --- /dev/null +++ b/lypi-contracts/src/main/java/cn/lypi/contracts/hook/BeforeTurnHookContext.java @@ -0,0 +1,20 @@ +package cn.lypi.contracts.hook; + +import cn.lypi.contracts.agent.TurnRequest; +import java.nio.file.Path; +import java.util.Objects; + +/** + * 表示 turn 执行前 hook 的上下文。 + */ +public record BeforeTurnHookContext( + TurnRequest request, + String turnId, + Path cwd +) { + public BeforeTurnHookContext { + request = Objects.requireNonNull(request, "request"); + turnId = Objects.requireNonNull(turnId, "turnId"); + cwd = Objects.requireNonNull(cwd, "cwd").toAbsolutePath().normalize(); + } +} diff --git a/lypi-contracts/src/main/java/cn/lypi/contracts/hook/BeforeTurnHookResult.java b/lypi-contracts/src/main/java/cn/lypi/contracts/hook/BeforeTurnHookResult.java new file mode 100644 index 00000000..b489acf0 --- /dev/null +++ b/lypi-contracts/src/main/java/cn/lypi/contracts/hook/BeforeTurnHookResult.java @@ -0,0 +1,23 @@ +package cn.lypi.contracts.hook; + +/** + * 表示 turn 执行前 hook 的决策结果。 + */ +public record BeforeTurnHookResult( + boolean blocked, + String message +) { + /** + * 返回允许继续执行的结果。 + */ + public static BeforeTurnHookResult allow() { + return new BeforeTurnHookResult(false, null); + } + + /** + * 返回阻断 turn 执行的结果。 + */ + public static BeforeTurnHookResult block(String message) { + return new BeforeTurnHookResult(true, message); + } +} diff --git a/lypi-contracts/src/main/java/cn/lypi/contracts/hook/DefaultToolHookRuntime.java b/lypi-contracts/src/main/java/cn/lypi/contracts/hook/DefaultToolHookRuntime.java new file mode 100644 index 00000000..c4dd448c --- /dev/null +++ b/lypi-contracts/src/main/java/cn/lypi/contracts/hook/DefaultToolHookRuntime.java @@ -0,0 +1,191 @@ +package cn.lypi.contracts.hook; + +import cn.lypi.contracts.event.EventBus; +import cn.lypi.contracts.event.HookEndEvent; +import cn.lypi.contracts.event.HookStartEvent; +import cn.lypi.contracts.tool.ToolResult; +import cn.lypi.contracts.tool.ToolUseContext; +import cn.lypi.contracts.tool.ToolUseRequest; +import java.time.Duration; +import java.time.Instant; +import java.util.List; +import java.util.Map; +import java.util.Objects; +import java.util.Optional; + +public final class DefaultToolHookRuntime implements ToolHookRuntime { + static final ToolHookRuntime NOOP = new DefaultToolHookRuntime(List.of()); + + private final List hooks; + private final EventBus eventBus; + + public DefaultToolHookRuntime(List hooks) { + this(hooks, null); + } + + public DefaultToolHookRuntime(List hooks, EventBus eventBus) { + this.hooks = List.copyOf(Objects.requireNonNull(hooks, "hooks")); + this.eventBus = eventBus; + } + + @Override + public BeforeToolHookResult beforeToolCall(BeforeToolHookContext context) { + BeforeToolHookContext nonNullContext = Objects.requireNonNull(context, "context"); + for (int index = 0; index < hooks.size(); index++) { + ToolHook hook = hooks.get(index); + HookRun run = start(nonNullContext.request(), nonNullContext.toolContext(), hook, HookPhase.BEFORE_TOOL_CALL, index); + try { + BeforeToolHookResult result = Objects.requireNonNull( + hook.beforeToolCall(nonNullContext), + "beforeToolCall result" + ); + if (result.blocked()) { + end(run, HookRunStatus.BLOCKED, result.message()); + return result; + } + end(run, HookRunStatus.SUCCEEDED, null); + } catch (RuntimeException exception) { + end(run, HookRunStatus.FAILED, exception.getMessage()); + throw exception; + } + } + return BeforeToolHookResult.allow(); + } + + @Override + public Optional> afterToolCall(AfterToolHookContext context) { + AfterToolHookContext currentContext = Objects.requireNonNull(context, "context"); + ToolResult currentResult = currentContext.result(); + boolean replaced = false; + for (int index = 0; index < hooks.size(); index++) { + ToolHook hook = hooks.get(index); + HookRun run = start(currentContext.request(), currentContext.toolContext(), hook, HookPhase.AFTER_TOOL_CALL, index); + try { + AfterToolHookResult hookResult = Objects.requireNonNull( + hook.afterToolCall(currentContext), + "afterToolCall result" + ); + Optional> replacement = hookResult.replacement(); + if (replacement.isPresent()) { + currentResult = replacement.orElseThrow(); + replaced = true; + currentContext = new AfterToolHookContext( + currentContext.request(), + currentContext.tool(), + currentContext.input(), + currentContext.toolContext(), + currentResult + ); + end(run, HookRunStatus.REPLACED, "工具结果已替换。"); + } else { + end(run, HookRunStatus.SUCCEEDED, null); + } + } catch (RuntimeException exception) { + end(run, HookRunStatus.FAILED, exception.getMessage()); + throw exception; + } + } + return replaced ? Optional.of(currentResult) : Optional.empty(); + } + + private HookRun start( + ToolUseRequest request, + ToolUseContext context, + ToolHook hook, + HookPhase phase, + int index + ) { + Instant startedAt = Instant.now(); + String hookRunId = hookRunId(request.toolUseId(), phase, index); + HookRun run = new HookRun( + safeText(context.sessionId(), "session_unknown"), + safeText(request.toolUseId(), "toolu_unknown"), + request.parentMessageId(), + turnId(context), + safeText(request.toolName(), "tool_unknown"), + hookRunId, + hookName(hook), + phase, + startedAt + ); + publish(new HookStartEvent( + run.sessionId(), + run.toolUseId(), + run.parentMessageId(), + run.turnId(), + run.toolName(), + run.hookRunId(), + run.hookName(), + run.phase(), + run.startedAt(), + startedAt + )); + return run; + } + + private void end(HookRun run, HookRunStatus status, String message) { + Instant endedAt = Instant.now(); + long durationMillis = Math.max(0L, Duration.between(run.startedAt(), endedAt).toMillis()); + publish(new HookEndEvent( + run.sessionId(), + run.toolUseId(), + run.parentMessageId(), + run.turnId(), + run.toolName(), + run.hookRunId(), + run.hookName(), + run.phase(), + status, + message, + run.startedAt(), + endedAt, + durationMillis, + endedAt + )); + } + + private void publish(cn.lypi.contracts.event.AgentEvent event) { + if (eventBus == null) { + return; + } + try { + eventBus.publish(event); + } catch (RuntimeException exception) { + // NOTE: hook 审计事件发布失败不得改变工具执行主链路。 + } + } + + private String turnId(ToolUseContext context) { + Object value = context.metadata() == null ? null : context.metadata().get("turnId"); + return value == null ? null : value.toString(); + } + + private String hookRunId(String toolUseId, HookPhase phase, int index) { + return "hook_" + safeText(toolUseId, "toolu_unknown") + "_" + phaseToken(phase) + "_" + index; + } + + private String hookName(ToolHook hook) { + String name = hook.name(); + return name == null || name.isBlank() ? hook.getClass().getName() : name; + } + + private String phaseToken(HookPhase phase) { + return phase == HookPhase.BEFORE_TOOL_CALL ? "before" : "after"; + } + + private String safeText(String value, String fallback) { + return value == null || value.isBlank() ? fallback : value; + } + + private record HookRun( + String sessionId, + String toolUseId, + String parentMessageId, + String turnId, + String toolName, + String hookRunId, + String hookName, + HookPhase phase, + Instant startedAt + ) {} +} diff --git a/lypi-contracts/src/main/java/cn/lypi/contracts/hook/DefaultTurnHookRuntime.java b/lypi-contracts/src/main/java/cn/lypi/contracts/hook/DefaultTurnHookRuntime.java new file mode 100644 index 00000000..b3ed0c3b --- /dev/null +++ b/lypi-contracts/src/main/java/cn/lypi/contracts/hook/DefaultTurnHookRuntime.java @@ -0,0 +1,40 @@ +package cn.lypi.contracts.hook; + +import java.util.List; +import java.util.Objects; + +public final class DefaultTurnHookRuntime implements TurnHookRuntime { + static final TurnHookRuntime NOOP = new DefaultTurnHookRuntime(List.of()); + + private final List hooks; + + public DefaultTurnHookRuntime(List hooks) { + this.hooks = List.copyOf(Objects.requireNonNull(hooks, "hooks")); + } + + @Override + public BeforeTurnHookResult beforeTurn(BeforeTurnHookContext context) { + BeforeTurnHookContext nonNullContext = Objects.requireNonNull(context, "context"); + for (TurnHook hook : hooks) { + BeforeTurnHookResult result = Objects.requireNonNull( + hook.beforeTurn(nonNullContext), + "beforeTurn result" + ); + if (result.blocked()) { + return result; + } + } + return BeforeTurnHookResult.allow(); + } + + @Override + public void afterTurn(AfterTurnHookContext context) { + AfterTurnHookContext nonNullContext = Objects.requireNonNull(context, "context"); + for (TurnHook hook : hooks) { + Objects.requireNonNull( + hook.afterTurn(nonNullContext), + "afterTurn result" + ); + } + } +} diff --git a/lypi-contracts/src/main/java/cn/lypi/contracts/hook/HookPhase.java b/lypi-contracts/src/main/java/cn/lypi/contracts/hook/HookPhase.java new file mode 100644 index 00000000..b499a787 --- /dev/null +++ b/lypi-contracts/src/main/java/cn/lypi/contracts/hook/HookPhase.java @@ -0,0 +1,6 @@ +package cn.lypi.contracts.hook; + +public enum HookPhase { + BEFORE_TOOL_CALL, + AFTER_TOOL_CALL +} diff --git a/lypi-contracts/src/main/java/cn/lypi/contracts/hook/HookRunStatus.java b/lypi-contracts/src/main/java/cn/lypi/contracts/hook/HookRunStatus.java new file mode 100644 index 00000000..1ea205c8 --- /dev/null +++ b/lypi-contracts/src/main/java/cn/lypi/contracts/hook/HookRunStatus.java @@ -0,0 +1,8 @@ +package cn.lypi.contracts.hook; + +public enum HookRunStatus { + SUCCEEDED, + BLOCKED, + REPLACED, + FAILED +} diff --git a/lypi-contracts/src/main/java/cn/lypi/contracts/hook/ToolHook.java b/lypi-contracts/src/main/java/cn/lypi/contracts/hook/ToolHook.java new file mode 100644 index 00000000..871517f7 --- /dev/null +++ b/lypi-contracts/src/main/java/cn/lypi/contracts/hook/ToolHook.java @@ -0,0 +1,74 @@ +package cn.lypi.contracts.hook; + +import java.util.Objects; + +public interface ToolHook { + /** + * 返回 hook 展示名称。 + */ + default String name() { + return getClass().getName(); + } + + /** + * 在工具执行前处理调用请求。 + * + * NOTE: 默认不阻断工具执行。 + */ + default BeforeToolHookResult beforeToolCall(BeforeToolHookContext context) { + Objects.requireNonNull(context, "context"); + return BeforeToolHookResult.allow(); + } + + /** + * 在工具执行后处理工具结果。 + * + * NOTE: 默认保留原始结果。 + */ + default AfterToolHookResult afterToolCall(AfterToolHookContext context) { + Objects.requireNonNull(context, "context"); + return AfterToolHookResult.keep(); + } + + /** + * 创建仅处理工具执行前阶段的 hook。 + */ + static ToolHook before(BeforeCallback callback) { + BeforeCallback nonNullCallback = Objects.requireNonNull(callback, "callback"); + return new ToolHook() { + @Override + public BeforeToolHookResult beforeToolCall(BeforeToolHookContext context) { + return Objects.requireNonNull(nonNullCallback.handle(context), "beforeToolCall result"); + } + }; + } + + /** + * 创建仅处理工具执行后阶段的 hook。 + */ + static ToolHook after(AfterCallback callback) { + AfterCallback nonNullCallback = Objects.requireNonNull(callback, "callback"); + return new ToolHook() { + @Override + public AfterToolHookResult afterToolCall(AfterToolHookContext context) { + return Objects.requireNonNull(nonNullCallback.handle(context), "afterToolCall result"); + } + }; + } + + @FunctionalInterface + interface BeforeCallback { + /** + * 处理工具执行前上下文并返回决策结果。 + */ + BeforeToolHookResult handle(BeforeToolHookContext context); + } + + @FunctionalInterface + interface AfterCallback { + /** + * 处理工具执行后上下文并返回结果处理决定。 + */ + AfterToolHookResult handle(AfterToolHookContext context); + } +} diff --git a/lypi-contracts/src/main/java/cn/lypi/contracts/hook/ToolHookInputSnapshots.java b/lypi-contracts/src/main/java/cn/lypi/contracts/hook/ToolHookInputSnapshots.java new file mode 100644 index 00000000..e7f5405d --- /dev/null +++ b/lypi-contracts/src/main/java/cn/lypi/contracts/hook/ToolHookInputSnapshots.java @@ -0,0 +1,39 @@ +package cn.lypi.contracts.hook; + +import java.util.ArrayList; +import java.util.Collections; +import java.util.LinkedHashMap; +import java.util.LinkedHashSet; +import java.util.List; +import java.util.Map; +import java.util.Set; + +final class ToolHookInputSnapshots { + private ToolHookInputSnapshots() { + } + + static Map snapshot(Map input) { + Map snapshot = new LinkedHashMap<>(); + input.forEach((key, value) -> snapshot.put(key, snapshotValue(value))); + return Collections.unmodifiableMap(snapshot); + } + + private static Object snapshotValue(Object value) { + if (value instanceof Map mapValue) { + Map snapshot = new LinkedHashMap<>(); + mapValue.forEach((key, nestedValue) -> snapshot.put(key, snapshotValue(nestedValue))); + return Collections.unmodifiableMap(snapshot); + } + if (value instanceof List listValue) { + List snapshot = new ArrayList<>(listValue.size()); + listValue.forEach(item -> snapshot.add(snapshotValue(item))); + return Collections.unmodifiableList(snapshot); + } + if (value instanceof Set setValue) { + Set snapshot = new LinkedHashSet<>(); + setValue.forEach(item -> snapshot.add(snapshotValue(item))); + return Collections.unmodifiableSet(snapshot); + } + return value; + } +} diff --git a/lypi-contracts/src/main/java/cn/lypi/contracts/hook/ToolHookRuntime.java b/lypi-contracts/src/main/java/cn/lypi/contracts/hook/ToolHookRuntime.java new file mode 100644 index 00000000..5c4f5de8 --- /dev/null +++ b/lypi-contracts/src/main/java/cn/lypi/contracts/hook/ToolHookRuntime.java @@ -0,0 +1,23 @@ +package cn.lypi.contracts.hook; + +import cn.lypi.contracts.tool.ToolResult; +import java.util.Optional; + +public interface ToolHookRuntime { + /** + * 顺序执行工具执行前 hook 并返回合成决策。 + */ + BeforeToolHookResult beforeToolCall(BeforeToolHookContext context); + + /** + * 顺序执行工具执行后 hook 并返回可选的最终替换结果。 + */ + Optional> afterToolCall(AfterToolHookContext context); + + /** + * 返回不执行任何 hook 的空实现。 + */ + static ToolHookRuntime noop() { + return DefaultToolHookRuntime.NOOP; + } +} diff --git a/lypi-contracts/src/main/java/cn/lypi/contracts/hook/TurnHook.java b/lypi-contracts/src/main/java/cn/lypi/contracts/hook/TurnHook.java new file mode 100644 index 00000000..f3e07efc --- /dev/null +++ b/lypi-contracts/src/main/java/cn/lypi/contracts/hook/TurnHook.java @@ -0,0 +1,74 @@ +package cn.lypi.contracts.hook; + +import java.util.Objects; + +public interface TurnHook { + /** + * 返回 hook 展示名称。 + */ + default String name() { + return getClass().getName(); + } + + /** + * 在 turn 执行前处理请求。 + * + * NOTE: 默认不阻断 turn 执行。 + */ + default BeforeTurnHookResult beforeTurn(BeforeTurnHookContext context) { + Objects.requireNonNull(context, "context"); + return BeforeTurnHookResult.allow(); + } + + /** + * 在 turn 执行后处理最终状态。 + * + * NOTE: 默认保留原始状态。 + */ + default AfterTurnHookResult afterTurn(AfterTurnHookContext context) { + Objects.requireNonNull(context, "context"); + return AfterTurnHookResult.keep(); + } + + /** + * 创建仅处理 turn 执行前阶段的 hook。 + */ + static TurnHook before(BeforeCallback callback) { + BeforeCallback nonNullCallback = Objects.requireNonNull(callback, "callback"); + return new TurnHook() { + @Override + public BeforeTurnHookResult beforeTurn(BeforeTurnHookContext context) { + return Objects.requireNonNull(nonNullCallback.handle(context), "beforeTurn result"); + } + }; + } + + /** + * 创建仅处理 turn 执行后阶段的 hook。 + */ + static TurnHook after(AfterCallback callback) { + AfterCallback nonNullCallback = Objects.requireNonNull(callback, "callback"); + return new TurnHook() { + @Override + public AfterTurnHookResult afterTurn(AfterTurnHookContext context) { + return Objects.requireNonNull(nonNullCallback.handle(context), "afterTurn result"); + } + }; + } + + @FunctionalInterface + interface BeforeCallback { + /** + * 处理 turn 执行前上下文并返回决策结果。 + */ + BeforeTurnHookResult handle(BeforeTurnHookContext context); + } + + @FunctionalInterface + interface AfterCallback { + /** + * 观察 turn 执行后上下文并返回处理完成结果。 + */ + AfterTurnHookResult handle(AfterTurnHookContext context); + } +} diff --git a/lypi-contracts/src/main/java/cn/lypi/contracts/hook/TurnHookRuntime.java b/lypi-contracts/src/main/java/cn/lypi/contracts/hook/TurnHookRuntime.java new file mode 100644 index 00000000..f023db07 --- /dev/null +++ b/lypi-contracts/src/main/java/cn/lypi/contracts/hook/TurnHookRuntime.java @@ -0,0 +1,20 @@ +package cn.lypi.contracts.hook; + +public interface TurnHookRuntime { + /** + * 顺序执行 turn 执行前 hook 并返回合成决策。 + */ + BeforeTurnHookResult beforeTurn(BeforeTurnHookContext context); + + /** + * 顺序执行 turn 执行后 hook。 + */ + void afterTurn(AfterTurnHookContext context); + + /** + * 返回不执行任何 hook 的空实现。 + */ + static TurnHookRuntime noop() { + return DefaultTurnHookRuntime.NOOP; + } +} diff --git a/lypi-contracts/src/test/java/cn/lypi/contracts/ContractSerializationTest.java b/lypi-contracts/src/test/java/cn/lypi/contracts/ContractSerializationTest.java index 01148a53..9a37e66b 100644 --- a/lypi-contracts/src/test/java/cn/lypi/contracts/ContractSerializationTest.java +++ b/lypi-contracts/src/test/java/cn/lypi/contracts/ContractSerializationTest.java @@ -22,6 +22,8 @@ import cn.lypi.contracts.error.ToolValidationException; import cn.lypi.contracts.event.AgentEvent; import cn.lypi.contracts.event.EventEnvelope; +import cn.lypi.contracts.event.HookEndEvent; +import cn.lypi.contracts.event.HookStartEvent; import cn.lypi.contracts.event.MessageBlockSnapshot; import cn.lypi.contracts.event.MessageDeltaEvent; import cn.lypi.contracts.event.MessageEndEvent; @@ -35,6 +37,8 @@ import cn.lypi.contracts.event.ToolStartEvent; import cn.lypi.contracts.event.TurnEndEvent; import cn.lypi.contracts.event.TurnStartEvent; +import cn.lypi.contracts.hook.HookPhase; +import cn.lypi.contracts.hook.HookRunStatus; import cn.lypi.contracts.model.AssistantStreamEvent; import cn.lypi.contracts.model.ModelSelection; import cn.lypi.contracts.model.ProviderRetryNotice; @@ -1111,6 +1115,62 @@ void toolEndEventRoundTripKeepsStatusSummaryRefAndTimingFields() throws Exceptio assertEquals(end.endedAt(), end.timestamp()); } + @Test + void hookStartEventRoundTripThroughAgentEvent() throws Exception { + Instant now = Instant.parse("2026-06-22T10:15:30Z"); + AgentEvent event = new HookStartEvent( + "session-1", + "toolu-1", + "msg-1", + "turn-1", + "Bash", + "hook_toolu-1_before_0", + "cn.lypi.TestHook", + HookPhase.BEFORE_TOOL_CALL, + now, + now + ); + + String json = mapper.writeValueAsString(event); + AgentEvent restored = mapper.readValue(json, AgentEvent.class); + + assertTrue(json.contains("\"type\":\"hook_start\"")); + HookStartEvent start = assertInstanceOf(HookStartEvent.class, restored); + assertEquals("session-1", start.sessionId()); + assertEquals("toolu-1", start.toolUseId()); + assertEquals(HookPhase.BEFORE_TOOL_CALL, start.phase()); + } + + @Test + void hookEndEventRoundTripThroughAgentEvent() throws Exception { + Instant startedAt = Instant.parse("2026-06-22T10:15:30Z"); + Instant endedAt = startedAt.plusMillis(25); + AgentEvent event = new HookEndEvent( + "session-1", + "toolu-1", + "msg-1", + "turn-1", + "Bash", + "hook_toolu-1_after_0", + "cn.lypi.TestHook", + HookPhase.AFTER_TOOL_CALL, + HookRunStatus.REPLACED, + "工具结果已替换。", + startedAt, + endedAt, + 25L, + endedAt + ); + + String json = mapper.writeValueAsString(event); + AgentEvent restored = mapper.readValue(json, AgentEvent.class); + + assertTrue(json.contains("\"type\":\"hook_end\"")); + HookEndEvent end = assertInstanceOf(HookEndEvent.class, restored); + assertEquals(HookRunStatus.REPLACED, end.status()); + assertEquals(25L, end.durationMillis()); + } + @Test void permissionRequestEventRoundTripContainsRenderableToolContext() throws Exception { PermissionDecision decision = new PermissionDecision( diff --git a/lypi-contracts/src/test/java/cn/lypi/contracts/hook/DefaultToolHookRuntimeTest.java b/lypi-contracts/src/test/java/cn/lypi/contracts/hook/DefaultToolHookRuntimeTest.java new file mode 100644 index 00000000..5e0c09da --- /dev/null +++ b/lypi-contracts/src/test/java/cn/lypi/contracts/hook/DefaultToolHookRuntimeTest.java @@ -0,0 +1,602 @@ +package cn.lypi.contracts.hook; + +import static org.junit.jupiter.api.Assertions.assertEquals; +import static org.junit.jupiter.api.Assertions.assertFalse; +import static org.junit.jupiter.api.Assertions.assertInstanceOf; +import static org.junit.jupiter.api.Assertions.assertNotSame; +import static org.junit.jupiter.api.Assertions.assertNull; +import static org.junit.jupiter.api.Assertions.assertThrows; +import static org.junit.jupiter.api.Assertions.assertTrue; +import static org.junit.jupiter.api.Assertions.fail; + +import cn.lypi.contracts.common.JsonSchema; +import cn.lypi.contracts.common.ProgressSink; +import cn.lypi.contracts.common.ValidationResult; +import cn.lypi.contracts.context.AgentMessage; +import cn.lypi.contracts.event.AgentEvent; +import cn.lypi.contracts.event.EventBus; +import cn.lypi.contracts.event.EventConsumer; +import cn.lypi.contracts.event.EventFilter; +import cn.lypi.contracts.event.EventSubscription; +import cn.lypi.contracts.event.HookEndEvent; +import cn.lypi.contracts.event.HookStartEvent; +import cn.lypi.contracts.security.PermissionBehavior; +import cn.lypi.contracts.security.PermissionDecision; +import cn.lypi.contracts.security.PermissionDecisionReason; +import cn.lypi.contracts.tool.InterruptBehavior; +import cn.lypi.contracts.tool.Tool; +import cn.lypi.contracts.tool.ToolResult; +import cn.lypi.contracts.tool.ToolUseContext; +import cn.lypi.contracts.tool.ToolUseRequest; +import java.nio.file.Path; +import java.util.ArrayList; +import java.util.Arrays; +import java.util.LinkedHashMap; +import java.util.LinkedHashSet; +import java.util.List; +import java.util.Map; +import java.util.Optional; +import java.util.Set; +import org.junit.jupiter.api.Test; + +class DefaultToolHookRuntimeTest { + @Test + void toolHookDefaultsToNoOpBeforeAndAfter() { + ToolHook hook = new ToolHook() { + }; + + BeforeToolHookResult beforeResult = hook.beforeToolCall(beforeContext()); + AfterToolHookResult afterResult = hook.afterToolCall(afterContext(result("original"))); + + assertFalse(beforeResult.blocked()); + assertNull(beforeResult.message()); + assertTrue(afterResult.replacement().isEmpty()); + } + + @Test + void noopRuntimeAllowsBeforeAndKeepsAfter() { + ToolHookRuntime runtime = ToolHookRuntime.noop(); + + BeforeToolHookResult beforeResult = runtime.beforeToolCall(beforeContext()); + Optional> afterResult = runtime.afterToolCall(afterContext(result("original"))); + + assertFalse(beforeResult.blocked()); + assertNull(beforeResult.message()); + assertTrue(afterResult.isEmpty()); + } + + @Test + void beforeStopsAtFirstBlockingHook() { + ToolHook first = ToolHook.before(context -> BeforeToolHookResult.block("blocked")); + ToolHook second = ToolHook.before(context -> { + fail("阻断后不应继续执行后续 before hook"); + return BeforeToolHookResult.allow(); + }); + + BeforeToolHookResult result = new DefaultToolHookRuntime(List.of(first, second)) + .beforeToolCall(beforeContext()); + + assertTrue(result.blocked()); + assertEquals("blocked", result.message()); + } + + @Test + void beforeAllowsWhenNoHookBlocks() { + ToolHook first = ToolHook.before(context -> BeforeToolHookResult.allow()); + ToolHook second = ToolHook.before(context -> BeforeToolHookResult.allow()); + + BeforeToolHookResult result = new DefaultToolHookRuntime(List.of(first, second)) + .beforeToolCall(beforeContext()); + + assertFalse(result.blocked()); + assertNull(result.message()); + } + + @Test + void afterAppliesResultReplacementInOrder() { + ToolHook first = ToolHook.after(context -> AfterToolHookResult.replace(result("first"))); + ToolHook second = ToolHook.after(context -> + AfterToolHookResult.replace(result(context.result().output() + "-second")) + ); + + ToolResult result = new DefaultToolHookRuntime(List.of(first, second)) + .afterToolCall(afterContext(result("original"))) + .orElseThrow(); + + assertEquals("first-second", result.output()); + } + + @Test + void afterReturnsEmptyWhenHooksKeepOriginal() { + ToolHook first = ToolHook.after(context -> AfterToolHookResult.keep()); + ToolHook second = ToolHook.after(context -> AfterToolHookResult.keep()); + + Optional> result = new DefaultToolHookRuntime(List.of(first, second)) + .afterToolCall(afterContext(result("original"))); + + assertTrue(result.isEmpty()); + } + + @Test + void runtimeDefensivelyCopiesHookList() { + List hooks = new ArrayList<>(); + hooks.add(ToolHook.before(context -> BeforeToolHookResult.allow())); + + DefaultToolHookRuntime runtime = new DefaultToolHookRuntime(hooks); + + hooks.clear(); + hooks.add(ToolHook.before(context -> BeforeToolHookResult.block("late-block"))); + + BeforeToolHookResult result = runtime.beforeToolCall(beforeContext()); + + assertFalse(result.blocked()); + } + + @Test + void beforePublishesStartAndBlockedEndEvent() { + CapturingEventBus events = new CapturingEventBus(); + ToolHook hook = ToolHook.before(context -> BeforeToolHookResult.block("denied")); + + BeforeToolHookResult result = new DefaultToolHookRuntime(List.of(hook), events) + .beforeToolCall(beforeContext()); + + assertTrue(result.blocked()); + assertEquals(2, events.events.size()); + HookStartEvent start = assertInstanceOf(HookStartEvent.class, events.events.get(0)); + HookEndEvent end = assertInstanceOf(HookEndEvent.class, events.events.get(1)); + assertEquals(HookPhase.BEFORE_TOOL_CALL, start.phase()); + assertEquals("ses_test", start.sessionId()); + assertEquals("toolu_test", start.toolUseId()); + assertEquals("msg_parent", start.parentMessageId()); + assertEquals("demo-tool", start.toolName()); + assertEquals(HookRunStatus.BLOCKED, end.status()); + assertEquals("denied", end.message()); + assertEquals(start.hookRunId(), end.hookRunId()); + } + + @Test + void afterPublishesReplacedEndEvent() { + CapturingEventBus events = new CapturingEventBus(); + ToolHook hook = ToolHook.after(context -> AfterToolHookResult.replace(result("rewritten"))); + + ToolResult result = new DefaultToolHookRuntime(List.of(hook), events) + .afterToolCall(afterContext(result("original"))) + .orElseThrow(); + + assertEquals("rewritten", result.output()); + assertEquals(2, events.events.size()); + HookEndEvent end = assertInstanceOf(HookEndEvent.class, events.events.get(1)); + assertEquals(HookPhase.AFTER_TOOL_CALL, end.phase()); + assertEquals(HookRunStatus.REPLACED, end.status()); + } + + @Test + void hookFailurePublishesFailedEndThenRethrows() { + CapturingEventBus events = new CapturingEventBus(); + ToolHook hook = ToolHook.before(context -> { + throw new IllegalStateException("boom"); + }); + + assertThrows(IllegalStateException.class, () -> new DefaultToolHookRuntime(List.of(hook), events) + .beforeToolCall(beforeContext())); + + assertEquals(2, events.events.size()); + HookEndEvent end = assertInstanceOf(HookEndEvent.class, events.events.get(1)); + assertEquals(HookRunStatus.FAILED, end.status()); + assertEquals("boom", end.message()); + } + + @Test + void eventPublishFailureDoesNotBlockHookExecution() { + EventBus failingEvents = new EventBus() { + @Override + public void publish(AgentEvent event) { + throw new IllegalStateException("event bus failed"); + } + + @Override + public EventSubscription subscribe(EventFilter filter, EventConsumer consumer) { + return () -> { + }; + } + }; + ToolHook hook = ToolHook.before(context -> BeforeToolHookResult.allow()); + + BeforeToolHookResult result = new DefaultToolHookRuntime(List.of(hook), failingEvents) + .beforeToolCall(beforeContext()); + + assertFalse(result.blocked()); + } + + @Test + void missingAuditMetadataDoesNotBlockHookExecution() { + CapturingEventBus events = new CapturingEventBus(); + ToolUseRequest request = new ToolUseRequest(null, null, Map.of("value", "demo"), null); + ToolUseContext context = new ToolUseContext(null, null, Path.of("."), Map.of()); + BeforeToolHookContext hookContext = new BeforeToolHookContext(request, tool(), Map.of("value", "demo"), context); + ToolHook hook = ToolHook.before(ignored -> BeforeToolHookResult.allow()); + + BeforeToolHookResult result = new DefaultToolHookRuntime(List.of(hook), events) + .beforeToolCall(hookContext); + + assertFalse(result.blocked()); + HookStartEvent start = assertInstanceOf(HookStartEvent.class, events.events.getFirst()); + assertEquals("session_unknown", start.sessionId()); + assertEquals("toolu_unknown", start.toolUseId()); + assertEquals("tool_unknown", start.toolName()); + } + + @Test + void beforeContextNormalizesInputForContextAndRequest() { + Map mutableInput = new LinkedHashMap<>(); + mutableInput.put("value", "demo"); + ToolUseRequest request = new ToolUseRequest("toolu_test", "demo-tool", mutableInput, "msg_parent"); + + BeforeToolHookContext context = new BeforeToolHookContext(request, tool(), mutableInput, toolContext()); + + mutableInput.put("late", "mutation"); + + assertEquals(Map.of("value", "demo"), context.input()); + assertEquals(context.input(), context.request().input()); + assertNotSame(mutableInput, context.input()); + assertThrows(UnsupportedOperationException.class, () -> context.input().put("extra", "x")); + assertThrows(UnsupportedOperationException.class, () -> context.request().input().put("extra", "x")); + } + + @Test + void beforeContextDeepCopiesNestedInputForContextAndRequest() { + List nestedList = new ArrayList<>(List.of("first", "second")); + Map nestedMap = new LinkedHashMap<>(); + nestedMap.put("flag", true); + nestedMap.put("items", nestedList); + Map mutableInput = new LinkedHashMap<>(); + mutableInput.put("nested", nestedMap); + + BeforeToolHookContext context = new BeforeToolHookContext( + new ToolUseRequest("toolu_test", "demo-tool", mutableInput, "msg_parent"), + tool(), + mutableInput, + toolContext() + ); + + nestedList.add("late-item"); + nestedMap.put("late-key", "late-value"); + + Map contextNestedMap = (Map) context.input().get("nested"); + List contextNestedList = (List) contextNestedMap.get("items"); + Map requestNestedMap = (Map) context.request().input().get("nested"); + List requestNestedList = (List) requestNestedMap.get("items"); + + assertEquals(List.of("first", "second"), contextNestedList); + assertEquals(List.of("first", "second"), requestNestedList); + assertFalse(contextNestedMap.containsKey("late-key")); + assertFalse(requestNestedMap.containsKey("late-key")); + assertNotSame(nestedMap, contextNestedMap); + assertNotSame(nestedList, contextNestedList); + assertThrows(UnsupportedOperationException.class, () -> putEntry(contextNestedMap, "extra", "x")); + assertThrows(UnsupportedOperationException.class, () -> addItem(contextNestedList, "x")); + assertThrows(UnsupportedOperationException.class, () -> putEntry(requestNestedMap, "extra", "x")); + assertThrows(UnsupportedOperationException.class, () -> addItem(requestNestedList, "x")); + } + + @Test + void beforeContextPreservesTopLevelNullValues() { + Map mutableInput = new LinkedHashMap<>(); + mutableInput.put("nullable", null); + + BeforeToolHookContext context = new BeforeToolHookContext( + new ToolUseRequest("toolu_test", "demo-tool", mutableInput, "msg_parent"), + tool(), + mutableInput, + toolContext() + ); + + assertTrue(context.input().containsKey("nullable")); + assertNull(context.input().get("nullable")); + assertTrue(context.request().input().containsKey("nullable")); + assertNull(context.request().input().get("nullable")); + assertThrows(UnsupportedOperationException.class, () -> context.input().put("extra", "x")); + assertThrows(UnsupportedOperationException.class, () -> context.request().input().put("extra", "x")); + } + + @Test + void beforeContextPreservesNullsInsideNestedMapAndList() { + List nestedList = new ArrayList<>(); + nestedList.add(null); + nestedList.add("second"); + Map nestedMap = new LinkedHashMap<>(); + nestedMap.put("nullable", null); + nestedMap.put("items", nestedList); + Map mutableInput = new LinkedHashMap<>(); + mutableInput.put("nested", nestedMap); + + BeforeToolHookContext context = new BeforeToolHookContext( + new ToolUseRequest("toolu_test", "demo-tool", mutableInput, "msg_parent"), + tool(), + mutableInput, + toolContext() + ); + + Map contextNestedMap = (Map) context.input().get("nested"); + List contextNestedList = (List) contextNestedMap.get("items"); + Map requestNestedMap = (Map) context.request().input().get("nested"); + List requestNestedList = (List) requestNestedMap.get("items"); + + assertTrue(contextNestedMap.containsKey("nullable")); + assertNull(contextNestedMap.get("nullable")); + assertEquals(Arrays.asList(null, "second"), contextNestedList); + assertTrue(requestNestedMap.containsKey("nullable")); + assertNull(requestNestedMap.get("nullable")); + assertEquals(Arrays.asList(null, "second"), requestNestedList); + assertThrows(UnsupportedOperationException.class, () -> putEntry(contextNestedMap, "extra", "x")); + assertThrows(UnsupportedOperationException.class, () -> addItem(contextNestedList, "x")); + assertThrows(UnsupportedOperationException.class, () -> putEntry(requestNestedMap, "extra", "x")); + assertThrows(UnsupportedOperationException.class, () -> addItem(requestNestedList, "x")); + } + + @Test + void afterContextNormalizesInputForContextAndRequest() { + Map mutableInput = new LinkedHashMap<>(); + mutableInput.put("value", "demo"); + ToolUseRequest request = new ToolUseRequest("toolu_test", "demo-tool", mutableInput, "msg_parent"); + + AfterToolHookContext context = new AfterToolHookContext( + request, + tool(), + mutableInput, + toolContext(), + result("original") + ); + + mutableInput.put("late", "mutation"); + + assertEquals(Map.of("value", "demo"), context.input()); + assertEquals(context.input(), context.request().input()); + assertNotSame(mutableInput, context.input()); + assertThrows(UnsupportedOperationException.class, () -> context.input().put("extra", "x")); + assertThrows(UnsupportedOperationException.class, () -> context.request().input().put("extra", "x")); + } + + @Test + void afterContextDeepCopiesNestedInputForContextAndRequest() { + List nestedList = new ArrayList<>(List.of("first", "second")); + Map nestedMap = new LinkedHashMap<>(); + nestedMap.put("flag", true); + nestedMap.put("items", nestedList); + Map mutableInput = new LinkedHashMap<>(); + mutableInput.put("nested", nestedMap); + + AfterToolHookContext context = new AfterToolHookContext( + new ToolUseRequest("toolu_test", "demo-tool", mutableInput, "msg_parent"), + tool(), + mutableInput, + toolContext(), + result("original") + ); + + nestedList.add("late-item"); + nestedMap.put("late-key", "late-value"); + + Map contextNestedMap = (Map) context.input().get("nested"); + List contextNestedList = (List) contextNestedMap.get("items"); + Map requestNestedMap = (Map) context.request().input().get("nested"); + List requestNestedList = (List) requestNestedMap.get("items"); + + assertEquals(List.of("first", "second"), contextNestedList); + assertEquals(List.of("first", "second"), requestNestedList); + assertFalse(contextNestedMap.containsKey("late-key")); + assertFalse(requestNestedMap.containsKey("late-key")); + assertNotSame(nestedMap, contextNestedMap); + assertNotSame(nestedList, contextNestedList); + assertThrows(UnsupportedOperationException.class, () -> putEntry(contextNestedMap, "extra", "x")); + assertThrows(UnsupportedOperationException.class, () -> addItem(contextNestedList, "x")); + assertThrows(UnsupportedOperationException.class, () -> putEntry(requestNestedMap, "extra", "x")); + assertThrows(UnsupportedOperationException.class, () -> addItem(requestNestedList, "x")); + } + + @Test + void afterContextPreservesNullsInsideNestedMapAndList() { + List nestedList = new ArrayList<>(); + nestedList.add(null); + nestedList.add("second"); + Map nestedMap = new LinkedHashMap<>(); + nestedMap.put("nullable", null); + nestedMap.put("items", nestedList); + Map mutableInput = new LinkedHashMap<>(); + mutableInput.put("nested", nestedMap); + + AfterToolHookContext context = new AfterToolHookContext( + new ToolUseRequest("toolu_test", "demo-tool", mutableInput, "msg_parent"), + tool(), + mutableInput, + toolContext(), + result("original") + ); + + Map contextNestedMap = (Map) context.input().get("nested"); + List contextNestedList = (List) contextNestedMap.get("items"); + Map requestNestedMap = (Map) context.request().input().get("nested"); + List requestNestedList = (List) requestNestedMap.get("items"); + + assertTrue(contextNestedMap.containsKey("nullable")); + assertNull(contextNestedMap.get("nullable")); + assertEquals(Arrays.asList(null, "second"), contextNestedList); + assertTrue(requestNestedMap.containsKey("nullable")); + assertNull(requestNestedMap.get("nullable")); + assertEquals(Arrays.asList(null, "second"), requestNestedList); + assertThrows(UnsupportedOperationException.class, () -> putEntry(contextNestedMap, "extra", "x")); + assertThrows(UnsupportedOperationException.class, () -> addItem(contextNestedList, "x")); + assertThrows(UnsupportedOperationException.class, () -> putEntry(requestNestedMap, "extra", "x")); + assertThrows(UnsupportedOperationException.class, () -> addItem(requestNestedList, "x")); + } + + @Test + void afterContextPreservesNullsInsideNestedSet() { + Set nestedSet = new LinkedHashSet<>(); + nestedSet.add(null); + nestedSet.add("value"); + Map mutableInput = new LinkedHashMap<>(); + mutableInput.put("set", nestedSet); + + AfterToolHookContext context = new AfterToolHookContext( + new ToolUseRequest("toolu_test", "demo-tool", mutableInput, "msg_parent"), + tool(), + mutableInput, + toolContext(), + result("original") + ); + + Set contextSet = (Set) context.input().get("set"); + Set requestSet = (Set) context.request().input().get("set"); + + nestedSet.add("late-value"); + + assertTrue(contextSet.contains(null)); + assertTrue(contextSet.contains("value")); + assertFalse(contextSet.contains("late-value")); + assertTrue(requestSet.contains(null)); + assertTrue(requestSet.contains("value")); + assertFalse(requestSet.contains("late-value")); + assertNotSame(nestedSet, contextSet); + assertThrows(UnsupportedOperationException.class, () -> addSetItem(contextSet, "x")); + assertThrows(UnsupportedOperationException.class, () -> addSetItem(requestSet, "x")); + } + + private BeforeToolHookContext beforeContext() { + return new BeforeToolHookContext( + request(), + tool(), + Map.of("value", "demo"), + toolContext() + ); + } + + private AfterToolHookContext afterContext(ToolResult result) { + return new AfterToolHookContext( + request(), + tool(), + Map.of("value", "demo"), + toolContext(), + result + ); + } + + private ToolUseRequest request() { + return new ToolUseRequest("toolu_test", "demo-tool", Map.of("value", "demo"), "msg_parent"); + } + + private ToolUseContext toolContext() { + return new ToolUseContext("ses_test", "msg_test", Path.of("."), Map.of("traceId", "trace-1")); + } + + @SuppressWarnings({ "rawtypes", "unchecked" }) + private void putEntry(Map map, Object key, Object value) { + ((Map) map).put(key, value); + } + + @SuppressWarnings({ "rawtypes", "unchecked" }) + private void addItem(List list, Object value) { + ((List) list).add(value); + } + + @SuppressWarnings({ "rawtypes", "unchecked" }) + private void addSetItem(Set set, Object value) { + ((Set) set).add(value); + } + + private Tool tool() { + return new DemoTool(); + } + + private ToolResult result(String output) { + return new ToolResult<>(output, false, List.of(), Optional.empty()); + } + + private static final class CapturingEventBus implements EventBus { + private final List events = new ArrayList<>(); + + @Override + public void publish(AgentEvent event) { + events.add(event); + } + + @Override + public EventSubscription subscribe(EventFilter filter, EventConsumer consumer) { + return () -> { + }; + } + } + + private static final class DemoTool implements Tool { + @Override + public String name() { + return "demo-tool"; + } + + @Override + public List aliases() { + return List.of(); + } + + @Override + public JsonSchema inputSchema() { + return new JsonSchema(Map.of("type", "object")); + } + + @Override + public ValidationResult validateInput(String input, ToolUseContext context) { + return new ValidationResult(true, List.of()); + } + + @Override + public PermissionDecision checkPermissions(String input, ToolUseContext context) { + return new PermissionDecision( + PermissionBehavior.ALLOW, + PermissionDecisionReason.TOOL_SPECIFIC, + "allowed", + Optional.empty(), + Map.of() + ); + } + + @Override + public ToolResult execute(String input, ToolUseContext context, ProgressSink progress) { + return new ToolResult<>(input, false, List.of(), Optional.empty()); + } + + @Override + public InterruptBehavior interruptBehavior() { + return InterruptBehavior.CANCEL; + } + + @Override + public boolean isReadOnly(String input) { + return true; + } + + @Override + public boolean isConcurrencySafe(String input) { + return true; + } + + @Override + public boolean isDestructive(String input) { + return false; + } + + @Override + public int maxResultSize() { + return 1024; + } + + @Override + public String renderForUser(String input) { + return input; + } + + @Override + public AgentMessage serializeForContext(String output) { + throw new UnsupportedOperationException("not needed for hook runtime tests"); + } + } +} diff --git a/lypi-contracts/src/test/java/cn/lypi/contracts/hook/DefaultTurnHookRuntimeTest.java b/lypi-contracts/src/test/java/cn/lypi/contracts/hook/DefaultTurnHookRuntimeTest.java new file mode 100644 index 00000000..0682730a --- /dev/null +++ b/lypi-contracts/src/test/java/cn/lypi/contracts/hook/DefaultTurnHookRuntimeTest.java @@ -0,0 +1,141 @@ +package cn.lypi.contracts.hook; + +import static org.junit.jupiter.api.Assertions.assertEquals; +import static org.junit.jupiter.api.Assertions.assertFalse; +import static org.junit.jupiter.api.Assertions.assertNull; +import static org.junit.jupiter.api.Assertions.assertThrows; +import static org.junit.jupiter.api.Assertions.assertTrue; +import static org.junit.jupiter.api.Assertions.fail; + +import cn.lypi.contracts.agent.TurnRequest; +import cn.lypi.contracts.agent.TurnState; +import cn.lypi.contracts.agent.TurnStatus; +import java.nio.file.Path; +import java.util.ArrayList; +import java.util.List; +import java.util.Optional; +import org.junit.jupiter.api.Test; + +class DefaultTurnHookRuntimeTest { + @Test + void turnHookDefaultsToNoOpBeforeAndAfter() { + TurnHook hook = new TurnHook() { + }; + + BeforeTurnHookResult beforeResult = hook.beforeTurn(beforeContext()); + AfterTurnHookResult afterResult = hook.afterTurn(afterContext(state(TurnStatus.COMPLETED, 0))); + + assertFalse(beforeResult.blocked()); + assertNull(beforeResult.message()); + assertEquals(AfterTurnHookResult.keep(), afterResult); + } + + @Test + void noopRuntimeAllowsBeforeAndKeepsAfter() { + TurnHookRuntime runtime = TurnHookRuntime.noop(); + + BeforeTurnHookResult beforeResult = runtime.beforeTurn(beforeContext()); + runtime.afterTurn(afterContext(state(TurnStatus.COMPLETED, 0))); + + assertFalse(beforeResult.blocked()); + assertNull(beforeResult.message()); + } + + @Test + void beforeStopsAtFirstBlockingHook() { + TurnHook first = TurnHook.before(context -> BeforeTurnHookResult.block("blocked")); + TurnHook second = TurnHook.before(context -> { + fail("阻断后不应继续执行后续 before turn hook"); + return BeforeTurnHookResult.allow(); + }); + + BeforeTurnHookResult result = new DefaultTurnHookRuntime(List.of(first, second)) + .beforeTurn(beforeContext()); + + assertTrue(result.blocked()); + assertEquals("blocked", result.message()); + } + + @Test + void beforeAllowsWhenNoHookBlocks() { + TurnHook first = TurnHook.before(context -> BeforeTurnHookResult.allow()); + TurnHook second = TurnHook.before(context -> BeforeTurnHookResult.allow()); + + BeforeTurnHookResult result = new DefaultTurnHookRuntime(List.of(first, second)) + .beforeTurn(beforeContext()); + + assertFalse(result.blocked()); + assertNull(result.message()); + } + + @Test + void afterRunsAllHooksInOrder() { + List observed = new ArrayList<>(); + TurnHook first = TurnHook.after(context -> { + observed.add(context.state().status()); + return AfterTurnHookResult.keep(); + }); + TurnHook second = TurnHook.after(context -> { + observed.add(context.state().status()); + return AfterTurnHookResult.keep(); + }); + + new DefaultTurnHookRuntime(List.of(first, second)) + .afterTurn(afterContext(state(TurnStatus.COMPLETED, 0))); + + assertEquals(List.of(TurnStatus.COMPLETED, TurnStatus.COMPLETED), observed); + } + + @Test + void afterReturnsEmptyWhenHooksKeepOriginal() { + TurnHook first = TurnHook.after(context -> AfterTurnHookResult.keep()); + TurnHook second = TurnHook.after(context -> AfterTurnHookResult.keep()); + + new DefaultTurnHookRuntime(List.of(first, second)) + .afterTurn(afterContext(state(TurnStatus.COMPLETED, 0))); + } + + @Test + void runtimeDefensivelyCopiesHookList() { + List hooks = new ArrayList<>(); + hooks.add(TurnHook.before(context -> BeforeTurnHookResult.allow())); + + DefaultTurnHookRuntime runtime = new DefaultTurnHookRuntime(hooks); + + hooks.clear(); + hooks.add(TurnHook.before(context -> BeforeTurnHookResult.block("late-block"))); + + BeforeTurnHookResult result = runtime.beforeTurn(beforeContext()); + + assertFalse(result.blocked()); + } + + @Test + void contextsRejectRequiredNulls() { + TurnRequest request = request(); + TurnState state = state(TurnStatus.COMPLETED, 0); + + assertThrows(NullPointerException.class, () -> new BeforeTurnHookContext(null, "turn-1", Path.of("."))); + assertThrows(NullPointerException.class, () -> new BeforeTurnHookContext(request, null, Path.of("."))); + assertThrows(NullPointerException.class, () -> new BeforeTurnHookContext(request, "turn-1", null)); + assertThrows(NullPointerException.class, () -> new AfterTurnHookContext(null, state, Path.of("."))); + assertThrows(NullPointerException.class, () -> new AfterTurnHookContext(request, null, Path.of("."))); + assertThrows(NullPointerException.class, () -> new AfterTurnHookContext(request, state, null)); + } + + private static BeforeTurnHookContext beforeContext() { + return new BeforeTurnHookContext(request(), "turn-1", Path.of("/tmp/project")); + } + + private static AfterTurnHookContext afterContext(TurnState state) { + return new AfterTurnHookContext(request(), state, Path.of("/tmp/project")); + } + + private static TurnRequest request() { + return new TurnRequest("session-1", "hello", Optional.empty(), () -> false); + } + + private static TurnState state(TurnStatus status, int toolRound) { + return new TurnState("turn-1", "session-1", null, List.of(), toolRound, status); + } +} diff --git a/lypi-tool/src/main/java/cn/lypi/tool/ToolHookExecutionInterceptor.java b/lypi-tool/src/main/java/cn/lypi/tool/ToolHookExecutionInterceptor.java new file mode 100644 index 00000000..8698f066 --- /dev/null +++ b/lypi-tool/src/main/java/cn/lypi/tool/ToolHookExecutionInterceptor.java @@ -0,0 +1,66 @@ +package cn.lypi.tool; + +import cn.lypi.contracts.hook.AfterToolHookContext; +import cn.lypi.contracts.hook.BeforeToolHookContext; +import cn.lypi.contracts.hook.BeforeToolHookResult; +import cn.lypi.contracts.hook.ToolHookRuntime; +import cn.lypi.contracts.tool.Tool; +import cn.lypi.contracts.tool.ToolResult; +import cn.lypi.contracts.tool.ToolUseContext; +import cn.lypi.contracts.tool.ToolUseRequest; +import java.util.Map; +import java.util.Objects; + +/** + * 将稳定 hook runtime 适配为工具执行拦截器。 + */ +public final class ToolHookExecutionInterceptor implements ToolExecutionInterceptor { + private final ToolHookRuntime runtime; + + /** + * 创建工具 hook 执行拦截器。 + * + * NOTE: 当 runtime 为 null 时回退到 no-op 实现,保持默认工具执行语义不变。 + */ + public ToolHookExecutionInterceptor(ToolHookRuntime runtime) { + this.runtime = runtime == null ? ToolHookRuntime.noop() : runtime; + } + + @Override + public BeforeResult beforeExecute(ToolUseRequest request, Tool, ?> tool, ToolUseContext context) { + ToolUseRequest nonNullRequest = Objects.requireNonNull(request, "request"); + Map canonicalInput = canonicalInput(nonNullRequest); + BeforeToolHookResult result = runtime.beforeToolCall(new BeforeToolHookContext( + nonNullRequest, + Objects.requireNonNull(tool, "tool"), + canonicalInput, + Objects.requireNonNull(context, "context") + )); + if (result != null && result.blocked()) { + return BeforeResult.block(result.message()); + } + return BeforeResult.allow(); + } + + @Override + public ToolResult afterExecute( + ToolUseRequest request, + Tool, ?> tool, + ToolUseContext context, + ToolResult result + ) { + ToolUseRequest nonNullRequest = Objects.requireNonNull(request, "request"); + Map canonicalInput = canonicalInput(nonNullRequest); + return runtime.afterToolCall(new AfterToolHookContext( + nonNullRequest, + Objects.requireNonNull(tool, "tool"), + canonicalInput, + Objects.requireNonNull(context, "context"), + Objects.requireNonNull(result, "result") + )).orElse(result); + } + + private Map canonicalInput(ToolUseRequest request) { + return request.input() == null ? Map.of() : request.input(); + } +} diff --git a/lypi-tool/src/test/java/cn/lypi/tool/ToolHookExecutionInterceptorTest.java b/lypi-tool/src/test/java/cn/lypi/tool/ToolHookExecutionInterceptorTest.java new file mode 100644 index 00000000..a9393572 --- /dev/null +++ b/lypi-tool/src/test/java/cn/lypi/tool/ToolHookExecutionInterceptorTest.java @@ -0,0 +1,204 @@ +package cn.lypi.tool; + +import static org.junit.jupiter.api.Assertions.assertEquals; +import static org.junit.jupiter.api.Assertions.assertFalse; +import static org.junit.jupiter.api.Assertions.assertNotSame; +import static org.junit.jupiter.api.Assertions.assertTrue; + +import cn.lypi.contracts.hook.AfterToolHookResult; +import cn.lypi.contracts.hook.BeforeToolHookResult; +import cn.lypi.contracts.hook.ToolHookRuntime; +import cn.lypi.contracts.security.PermissionMode; +import cn.lypi.contracts.tool.ToolResult; +import cn.lypi.contracts.tool.ToolUseContext; +import cn.lypi.contracts.tool.ToolUseRequest; +import java.util.LinkedHashMap; +import java.util.List; +import java.util.Map; +import java.util.concurrent.atomic.AtomicReference; +import org.junit.jupiter.api.Test; + +class ToolHookExecutionInterceptorTest { + @Test + void beforeHookBlocksThroughInterceptor() { + ToolHookRuntime runtime = new ToolHookRuntime() { + @Override + public BeforeToolHookResult beforeToolCall(cn.lypi.contracts.hook.BeforeToolHookContext context) { + return BeforeToolHookResult.block("hook denied"); + } + + @Override + public java.util.Optional> afterToolCall(cn.lypi.contracts.hook.AfterToolHookContext context) { + return java.util.Optional.empty(); + } + }; + ToolHookExecutionInterceptor interceptor = new ToolHookExecutionInterceptor(runtime); + + ToolExecutionInterceptor.BeforeResult result = interceptor.beforeExecute( + request(Map.of()), + tool(), + context() + ); + + assertTrue(result.blocked()); + assertEquals("hook denied", result.message()); + } + + @Test + void afterHookCanReplaceToolResult() { + ToolHookRuntime runtime = new ToolHookRuntime() { + @Override + public BeforeToolHookResult beforeToolCall(cn.lypi.contracts.hook.BeforeToolHookContext context) { + return BeforeToolHookResult.allow(); + } + + @Override + public java.util.Optional> afterToolCall(cn.lypi.contracts.hook.AfterToolHookContext context) { + return AfterToolHookResult.replace(TestTools.result(context.request().toolUseId(), "rewritten", false)) + .replacement(); + } + }; + ToolHookExecutionInterceptor interceptor = new ToolHookExecutionInterceptor(runtime); + + ToolResult result = interceptor.afterExecute( + request(Map.of()), + tool(), + context(), + TestTools.result("toolu_1", "original", false) + ); + + assertEquals("rewritten", result.output()); + } + + @Test + void nullRuntimeFallsBackToNoop() { + ToolHookExecutionInterceptor interceptor = new ToolHookExecutionInterceptor(null); + ToolUseRequest request = request(Map.of("text", "hello")); + ToolResult original = TestTools.result(request.toolUseId(), "original", false); + + ToolExecutionInterceptor.BeforeResult beforeResult = interceptor.beforeExecute(request, tool(), context()); + ToolResult afterResult = interceptor.afterExecute(request, tool(), context(), original); + + assertFalse(beforeResult.blocked()); + assertEquals("", beforeResult.message()); + assertSameResult(original, afterResult); + } + + @Test + void contextsSeeCanonicalSnapshotInput() { + AtomicReference> beforeInput = new AtomicReference<>(); + AtomicReference> beforeRequestInput = new AtomicReference<>(); + AtomicReference> afterInput = new AtomicReference<>(); + AtomicReference> afterRequestInput = new AtomicReference<>(); + ToolHookRuntime runtime = new ToolHookRuntime() { + @Override + public BeforeToolHookResult beforeToolCall(cn.lypi.contracts.hook.BeforeToolHookContext context) { + beforeInput.set(context.input()); + beforeRequestInput.set(context.request().input()); + return BeforeToolHookResult.allow(); + } + + @Override + public java.util.Optional> afterToolCall(cn.lypi.contracts.hook.AfterToolHookContext context) { + afterInput.set(context.input()); + afterRequestInput.set(context.request().input()); + return java.util.Optional.empty(); + } + }; + ToolHookExecutionInterceptor interceptor = new ToolHookExecutionInterceptor(runtime); + Map mutableInput = new LinkedHashMap<>(); + Map nested = new LinkedHashMap<>(); + nested.put("flag", true); + mutableInput.put("nested", nested); + ToolUseRequest request = request(mutableInput); + + interceptor.beforeExecute(request, tool(), context()); + interceptor.afterExecute(request, tool(), context(), TestTools.result(request.toolUseId(), "original", false)); + nested.put("late", "mutation"); + mutableInput.put("newKey", "newValue"); + + assertEquals(beforeInput.get(), beforeRequestInput.get()); + assertEquals(afterInput.get(), afterRequestInput.get()); + assertEquals(beforeInput.get(), afterInput.get()); + assertNotSame(mutableInput, beforeInput.get()); + assertNotSame(mutableInput, beforeRequestInput.get()); + assertFalse(beforeInput.get().containsKey("newKey")); + assertFalse(beforeRequestInput.get().containsKey("newKey")); + Map beforeNested = (Map) beforeInput.get().get("nested"); + Map afterNested = (Map) afterInput.get().get("nested"); + assertFalse(beforeNested.containsKey("late")); + assertFalse(afterNested.containsKey("late")); + } + + @Test + void beforeHookSeesEmptyMapWhenRequestInputIsNull() { + AtomicReference> beforeInput = new AtomicReference<>(); + AtomicReference> beforeRequestInput = new AtomicReference<>(); + ToolHookRuntime runtime = new ToolHookRuntime() { + @Override + public BeforeToolHookResult beforeToolCall(cn.lypi.contracts.hook.BeforeToolHookContext context) { + beforeInput.set(context.input()); + beforeRequestInput.set(context.request().input()); + return BeforeToolHookResult.allow(); + } + + @Override + public java.util.Optional> afterToolCall(cn.lypi.contracts.hook.AfterToolHookContext context) { + return java.util.Optional.empty(); + } + }; + ToolHookExecutionInterceptor interceptor = new ToolHookExecutionInterceptor(runtime); + + ToolExecutionInterceptor.BeforeResult result = interceptor.beforeExecute(request(null), tool(), context()); + + assertFalse(result.blocked()); + assertEquals(Map.of(), beforeInput.get()); + assertEquals(Map.of(), beforeRequestInput.get()); + } + + @Test + void afterHookSeesEmptyMapWhenRequestInputIsNull() { + AtomicReference> afterInput = new AtomicReference<>(); + AtomicReference> afterRequestInput = new AtomicReference<>(); + ToolHookRuntime runtime = new ToolHookRuntime() { + @Override + public BeforeToolHookResult beforeToolCall(cn.lypi.contracts.hook.BeforeToolHookContext context) { + return BeforeToolHookResult.allow(); + } + + @Override + public java.util.Optional> afterToolCall(cn.lypi.contracts.hook.AfterToolHookContext context) { + afterInput.set(context.input()); + afterRequestInput.set(context.request().input()); + return java.util.Optional.empty(); + } + }; + ToolHookExecutionInterceptor interceptor = new ToolHookExecutionInterceptor(runtime); + ToolResult original = TestTools.result("toolu_1", "original", false); + + ToolResult result = interceptor.afterExecute(request(null), tool(), context(), original); + + assertSameResult(original, result); + assertEquals(Map.of(), afterInput.get()); + assertEquals(Map.of(), afterRequestInput.get()); + } + + private ToolUseRequest request(Map input) { + return new ToolUseRequest("toolu_1", "echo", input, "msg_1"); + } + + private cn.lypi.contracts.tool.Tool, String> tool() { + return TestTools.echo("echo", List.of(), true, true, false); + } + + private ToolUseContext context() { + return TestTools.toolContext(PermissionMode.DEFAULT_EXECUTE); + } + + private void assertSameResult(ToolResult expected, ToolResult actual) { + assertEquals(expected.output(), actual.output()); + assertEquals(expected.isError(), actual.isError()); + assertEquals(expected.newMessages(), actual.newMessages()); + assertEquals(expected.replacement(), actual.replacement()); + } +} diff --git a/lypi-transport-tui/src/test/java/cn/lypi/transport/tui/TuiEventReducerTest.java b/lypi-transport-tui/src/test/java/cn/lypi/transport/tui/TuiEventReducerTest.java index 8aea0d55..c58405eb 100644 --- a/lypi-transport-tui/src/test/java/cn/lypi/transport/tui/TuiEventReducerTest.java +++ b/lypi-transport-tui/src/test/java/cn/lypi/transport/tui/TuiEventReducerTest.java @@ -16,6 +16,8 @@ import cn.lypi.contracts.event.CompactEndEvent; import cn.lypi.contracts.event.CompactStartEvent; import cn.lypi.contracts.event.ErrorEvent; +import cn.lypi.contracts.event.HookEndEvent; +import cn.lypi.contracts.event.HookStartEvent; import cn.lypi.contracts.event.InterruptEvent; import cn.lypi.contracts.event.MessageBlockSnapshot; import cn.lypi.contracts.event.MessageDeltaEvent; @@ -45,6 +47,8 @@ import cn.lypi.contracts.event.ToolStartEvent; import cn.lypi.contracts.event.TurnEndEvent; import cn.lypi.contracts.event.TurnStartEvent; +import cn.lypi.contracts.hook.HookPhase; +import cn.lypi.contracts.hook.HookRunStatus; import cn.lypi.contracts.security.PermissionBehavior; import cn.lypi.contracts.security.PermissionDecision; import cn.lypi.contracts.security.PermissionDecisionReason; @@ -401,6 +405,43 @@ void runtimeEventsUpdateEphemeralRuntimeLineWithoutAddingBlocks() { assertEquals(0, reducer.view().blocks().size()); } + @Test + void hookLifecycleEventsDoNotChangeTuiView() { + TuiEventReducer reducer = new TuiEventReducer(); + TuiViewModel before = reducer.view(); + + reducer.reduce(new HookStartEvent( + "ses_1", + "toolu_1", + "msg_1", + "turn_1", + "Bash", + "hook_toolu_1_before_0", + "cn.lypi.TestHook", + HookPhase.BEFORE_TOOL_CALL, + NOW, + NOW + )); + TuiViewModel after = reducer.reduce(new HookEndEvent( + "ses_1", + "toolu_1", + "msg_1", + "turn_1", + "Bash", + "hook_toolu_1_before_0", + "cn.lypi.TestHook", + HookPhase.BEFORE_TOOL_CALL, + HookRunStatus.SUCCEEDED, + null, + NOW, + NOW.plusMillis(1), + 1L, + NOW.plusMillis(1) + )); + + assertEquals(before, after); + } + @Test void turnRuntimeLineShowsWorkingElapsedTime() { TuiEventReducer reducer = TuiEventReducer.withRuntimeState(TestRuntimeStates.basic("ses_1"));