This is an automated email from the ASF dual-hosted git repository.
wenjin272 pushed a commit to branch main
in repository https://gitbox.apache.org/repos/asf/flink-agents.git
The following commit(s) were added to refs/heads/main by this push:
new 04bf85163 [plan][runtime] Do not retry or persist interrupted chat
calls on cancellation (#1071)
04bf85163 is described below
commit 04bf85163c733e4006e47a75685897fd5820f501
Author: Ashfaq <[email protected]>
AuthorDate: Wed Sep 9 13:44:50 2026 +0530
[plan][runtime] Do not retry or persist interrupted chat calls on
cancellation (#1071)
Generated-by: Claude Code 2.1.226 (Claude Sonnet 5)
---
.../agents/plan/actions/ChatModelInvoker.java | 13 +-
.../agents/plan/actions/ChatModelInvokerTest.java | 161 +++++++++++++++++++++
.../agents/runtime/context/RunnerContextImpl.java | 12 ++
.../RunnerContextImplDurableExecuteTest.java | 55 +++++++
4 files changed, 240 insertions(+), 1 deletion(-)
diff --git
a/plan/src/main/java/org/apache/flink/agents/plan/actions/ChatModelInvoker.java
b/plan/src/main/java/org/apache/flink/agents/plan/actions/ChatModelInvoker.java
index 4f49feac7..5d2193d5d 100644
---
a/plan/src/main/java/org/apache/flink/agents/plan/actions/ChatModelInvoker.java
+++
b/plan/src/main/java/org/apache/flink/agents/plan/actions/ChatModelInvoker.java
@@ -185,6 +185,12 @@ final class ChatModelInvoker {
}
return new ChatAttemptResult(
model, chatModel, response, actualRetryCount,
totalWaitTimeSec);
+ } catch (InterruptedException e) {
+ // A cancellation signal, not a model failure: restore the
interrupt status and
+ // propagate immediately so task shutdown isn't delayed by
retry backoff or an
+ // extra model call, regardless of the configured
error-handling strategy.
+ Thread.currentThread().interrupt();
+ throw e;
} catch (Exception e) {
if (strategy == Agent.ErrorHandlingStrategy.RETRY && attempt <
numRetries) {
actualRetryCount = attempt + 1;
@@ -197,7 +203,12 @@ final class ChatModelInvoker {
numRetries,
currentWaitSec);
if (currentWaitSec > 0) {
- Thread.sleep(currentWaitSec * 1000L);
+ try {
+ Thread.sleep(currentWaitSec * 1000L);
+ } catch (InterruptedException ie) {
+ Thread.currentThread().interrupt();
+ throw ie;
+ }
totalWaitTimeSec += currentWaitSec;
}
continue;
diff --git
a/plan/src/test/java/org/apache/flink/agents/plan/actions/ChatModelInvokerTest.java
b/plan/src/test/java/org/apache/flink/agents/plan/actions/ChatModelInvokerTest.java
new file mode 100644
index 000000000..9b403e351
--- /dev/null
+++
b/plan/src/test/java/org/apache/flink/agents/plan/actions/ChatModelInvokerTest.java
@@ -0,0 +1,161 @@
+/*
+ * Licensed to the Apache Software Foundation (ASF) under one
+ * or more contributor license agreements. See the NOTICE file
+ * distributed with this work for additional information
+ * regarding copyright ownership. The ASF licenses this file
+ * to you under the Apache License, Version 2.0 (the
+ * "License"); you may not use this file except in compliance
+ * with the License. You may obtain a copy of the License at
+ *
+ * http://www.apache.org/licenses/LICENSE-2.0
+ *
+ * Unless required by applicable law or agreed to in writing, software
+ * distributed under the License is distributed on an "AS IS" BASIS,
+ * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+ * See the License for the specific language governing permissions and
+ * limitations under the License.
+ */
+package org.apache.flink.agents.plan.actions;
+
+import org.apache.flink.agents.api.agents.Agent;
+import org.apache.flink.agents.api.agents.AgentExecutionOptions;
+import org.apache.flink.agents.api.chat.model.BaseChatModelSetup;
+import org.apache.flink.agents.api.configuration.ReadableConfiguration;
+import org.apache.flink.agents.api.context.RunnerContext;
+import org.apache.flink.agents.api.metrics.FlinkAgentsMetricGroup;
+import org.apache.flink.agents.api.resource.ResourceType;
+import org.junit.jupiter.api.AfterEach;
+import org.junit.jupiter.api.Test;
+
+import java.util.List;
+import java.util.Map;
+import java.util.UUID;
+
+import static org.junit.jupiter.api.Assertions.assertThrows;
+import static org.junit.jupiter.api.Assertions.assertTrue;
+import static org.mockito.ArgumentMatchers.any;
+import static org.mockito.Mockito.mock;
+import static org.mockito.Mockito.times;
+import static org.mockito.Mockito.verify;
+import static org.mockito.Mockito.when;
+
+/** Tests for {@link ChatModelInvoker}. */
+class ChatModelInvokerTest {
+
+ @AfterEach
+ void clearInterruptStatus() {
+ // Prevents a leftover interrupt flag (e.g. if the assertion below the
interruption test
+ // ever fails before consuming it) from failing an unrelated later
test's real Thread.sleep
+ // backoff with a spurious InterruptedException.
+ Thread.interrupted();
+ }
+
+ @Test
+ void testChatWithRetriesDoesNotRetryOnInterruption() throws Exception {
+ RunnerContext ctx = mock(RunnerContext.class);
+ BaseChatModelSetup chatModel = mock(BaseChatModelSetup.class);
+ ReadableConfiguration config = mock(ReadableConfiguration.class);
+ when(ctx.getConfig()).thenReturn(config);
+ when(config.get(AgentExecutionOptions.CHAT_ASYNC)).thenReturn(false);
+ when(ctx.getResource("test-model",
ResourceType.CHAT_MODEL)).thenReturn(chatModel);
+
when(ctx.getActionMetricGroup()).thenReturn(mock(FlinkAgentsMetricGroup.class));
+ when(ctx.durableExecute(any())).thenThrow(new
InterruptedException("cancelled"));
+
+ // Clear any interrupt status left over from a previous test before
asserting on it below.
+ Thread.interrupted();
+
+ assertThrows(
+ InterruptedException.class,
+ () ->
+ ChatModelInvoker.chatWithRetries(
+ UUID.randomUUID(),
+ "test-model",
+ "durable-call-id",
+ List.of(),
+ Map.of(),
+ null,
+ ctx,
+ Agent.ErrorHandlingStrategy.RETRY,
+ 3,
+ 0));
+
+ // Only the first attempt should have run: retry backoff must not
consume more attempts
+ // after a cancellation interrupts the call.
+ verify(ctx, times(1)).durableExecute(any());
+ assertTrue(Thread.interrupted(), "interrupt status should be restored
on the thread");
+ }
+
+ @Test
+ void testChatWithRetriesRestoresInterruptFlagFromRetryBackoffSleep()
throws Exception {
+ RunnerContext ctx = mock(RunnerContext.class);
+ BaseChatModelSetup chatModel = mock(BaseChatModelSetup.class);
+ ReadableConfiguration config = mock(ReadableConfiguration.class);
+ when(ctx.getConfig()).thenReturn(config);
+ when(config.get(AgentExecutionOptions.CHAT_ASYNC)).thenReturn(false);
+ when(ctx.getResource("test-model",
ResourceType.CHAT_MODEL)).thenReturn(chatModel);
+
when(ctx.getActionMetricGroup()).thenReturn(mock(FlinkAgentsMetricGroup.class));
+ // The interrupt fires from within the retry backoff's Thread.sleep,
not from the call
+ // itself: set the flag first so Thread.sleep throws immediately, at
no wall-clock cost.
+ when(ctx.durableExecute(any()))
+ .thenAnswer(
+ invocation -> {
+ Thread.currentThread().interrupt();
+ throw new RuntimeException("transient failure");
+ });
+
+ assertThrows(
+ InterruptedException.class,
+ () ->
+ ChatModelInvoker.chatWithRetries(
+ UUID.randomUUID(),
+ "test-model",
+ "durable-call-id",
+ List.of(),
+ Map.of(),
+ null,
+ ctx,
+ Agent.ErrorHandlingStrategy.RETRY,
+ 1,
+ 1));
+
+ // Only the first attempt should have run: the sleep before the retry
throws before a
+ // second call is made.
+ verify(ctx, times(1)).durableExecute(any());
+ // This is the only assertion that distinguishes the fix from the
pre-fix code: the
+ // InterruptedException/times(1) shape above passes either way, since
Thread.sleep still
+ // aborts the retry loop on both. Only the restored flag proves the
backoff sleep's catch
+ // block restores it instead of leaving it cleared.
+ assertTrue(Thread.interrupted(), "interrupt status should be restored
on the thread");
+ }
+
+ @Test
+ void testChatWithRetriesRetriesOnOrdinaryFailure() throws Exception {
+ RunnerContext ctx = mock(RunnerContext.class);
+ BaseChatModelSetup chatModel = mock(BaseChatModelSetup.class);
+ ReadableConfiguration config = mock(ReadableConfiguration.class);
+ when(ctx.getConfig()).thenReturn(config);
+ when(config.get(AgentExecutionOptions.CHAT_ASYNC)).thenReturn(false);
+ when(ctx.getResource("test-model",
ResourceType.CHAT_MODEL)).thenReturn(chatModel);
+
when(ctx.getActionMetricGroup()).thenReturn(mock(FlinkAgentsMetricGroup.class));
+ when(ctx.durableExecute(any())).thenThrow(new
RuntimeException("transient failure"));
+
+ assertThrows(
+ ChatModelInvoker.ChatAttemptFailed.class,
+ () ->
+ ChatModelInvoker.chatWithRetries(
+ UUID.randomUUID(),
+ "test-model",
+ "durable-call-id",
+ List.of(),
+ Map.of(),
+ null,
+ ctx,
+ Agent.ErrorHandlingStrategy.RETRY,
+ 2,
+ 0));
+
+ // An ordinary failure must still consume the full retry budget
(initial attempt + 2
+ // retries), confirming the interruption fix doesn't disturb normal
retry behavior.
+ verify(ctx, times(3)).durableExecute(any());
+ }
+}
diff --git
a/runtime/src/main/java/org/apache/flink/agents/runtime/context/RunnerContextImpl.java
b/runtime/src/main/java/org/apache/flink/agents/runtime/context/RunnerContextImpl.java
index dcfc34ed0..a40e56b03 100644
---
a/runtime/src/main/java/org/apache/flink/agents/runtime/context/RunnerContextImpl.java
+++
b/runtime/src/main/java/org/apache/flink/agents/runtime/context/RunnerContextImpl.java
@@ -576,6 +576,12 @@ public class RunnerContextImpl implements RunnerContext,
ExecutionReporter {
Exception exception = null;
try {
result = executionCallable.call();
+ } catch (InterruptedException e) {
+ // A cancellation signal, not a genuine call failure: leave the
durable slot
+ // unfinished so recovery re-executes or reconciles the call
instead of replaying a
+ // stale interruption as a completed success or failure.
+ Thread.currentThread().interrupt();
+ throw e;
} catch (Exception e) {
exception = e;
}
@@ -939,6 +945,12 @@ public class RunnerContextImpl implements RunnerContext,
ExecutionReporter {
Exception exception = null;
try {
result = callSupplier.call();
+ } catch (InterruptedException e) {
+ // A cancellation signal, not a genuine call failure: leave the
pending call
+ // unfinalized so recovery re-executes or reconciles it instead of
replaying a stale
+ // interruption as a completed success or failure.
+ Thread.currentThread().interrupt();
+ throw e;
} catch (Exception e) {
exception = e;
}
diff --git
a/runtime/src/test/java/org/apache/flink/agents/runtime/context/RunnerContextImplDurableExecuteTest.java
b/runtime/src/test/java/org/apache/flink/agents/runtime/context/RunnerContextImplDurableExecuteTest.java
index 7d4cfde21..816b874d5 100644
---
a/runtime/src/test/java/org/apache/flink/agents/runtime/context/RunnerContextImplDurableExecuteTest.java
+++
b/runtime/src/test/java/org/apache/flink/agents/runtime/context/RunnerContextImplDurableExecuteTest.java
@@ -26,6 +26,7 @@ import
org.apache.flink.agents.runtime.actionstate.ActionState;
import org.apache.flink.agents.runtime.actionstate.CallResult;
import org.apache.flink.agents.runtime.metrics.FlinkAgentsMetricGroupImpl;
import org.apache.flink.runtime.metrics.groups.UnregisteredMetricGroups;
+import org.junit.jupiter.api.AfterEach;
import org.junit.jupiter.api.BeforeEach;
import org.junit.jupiter.api.Test;
@@ -53,6 +54,13 @@ class RunnerContextImplDurableExecuteTest {
lastPersistedState = null;
}
+ @AfterEach
+ void clearInterruptStatus() {
+ // Prevents a leftover interrupt flag (e.g. if an assertion below one
of the interruption
+ // tests ever fails before consuming it) from leaking into an
unrelated later test.
+ Thread.interrupted();
+ }
+
@Test
void testDurableExecuteLegacyCall() throws Exception {
RunnerContextImpl context = createContext(new ActionState(null));
@@ -72,6 +80,53 @@ class RunnerContextImplDurableExecuteTest {
assertSame(context.getDurableExecutionContext().getActionState(),
lastPersistedState);
}
+ @Test
+ void testDurableExecuteCompletionOnlyDoesNotPersistInterruption() {
+ RunnerContextImpl context = createContext(new ActionState(null));
+ TestDurableCallable<String> callable =
+ new TestDurableCallable<>(
+ "legacy-call",
+ String.class,
+ () -> {
+ throw new InterruptedException("cancelled");
+ });
+
+ // Clear any interrupt status left over from a previous test before
asserting on it below.
+ Thread.interrupted();
+
+ assertThrows(InterruptedException.class, () ->
context.durableExecute(callable));
+
+ assertTrue(Thread.interrupted(), "interrupt status should be restored
on the thread");
+ // A cancellation must not be finalized as a durable success or
failure: the slot stays
+ // unfinished so recovery re-executes the call instead of replaying a
stale interruption.
+ assertEquals(0, persistCallCount.get());
+ assertEquals(0,
context.getDurableExecutionContext().getActionState().getCallResultCount());
+ }
+
+ @Test
+ void
testDurableExecuteCompletionOnlyReExecutesPendingSlotDoesNotPersistInterruption()
{
+ ActionState actionState = new ActionState(null);
+ actionState.addCallResult(CallResult.pending("tool-call", ""));
+ RunnerContextImpl context = createContext(actionState);
+ TestDurableCallable<String> callable =
+ new TestDurableCallable<>(
+ "tool-call",
+ String.class,
+ () -> {
+ throw new InterruptedException("cancelled");
+ });
+
+ Thread.interrupted();
+
+ assertThrows(InterruptedException.class, () ->
context.durableExecute(callable));
+
+ assertTrue(Thread.interrupted(), "interrupt status should be restored
on the thread");
+ assertEquals(0, persistCallCount.get());
+ CallResult pending =
+
context.getDurableExecutionContext().getActionState().getCallResults().get(0);
+ assertTrue(pending.isPending(), "interrupted pending slot should
remain unfinalized");
+ }
+
@Test
void testDurableExecuteReconcilableSuccessCall() throws Exception {
RunnerContextImpl context = createContext(new ActionState(null));