Skip to content

Commit 732cad8

Browse files
committed
Add LLM retry and timeout handling to agent loop
1 parent 0043e6d commit 732cad8

4 files changed

Lines changed: 425 additions & 2 deletions

File tree

README.md

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -248,6 +248,8 @@ The next minor release focuses on interview-ready onboarding and stronger runtim
248248
- [x] Add an Agent loop architecture diagram to this README (message/tool/confirmation flow).
249249
- [x] Add core tests for `AgentLoop` and `SkillLoader` to balance existing provider-heavy coverage.
250250
- [ ] Add clearer memory/context guidance and APIs for session-level context management.
251+
- [ ] Add an optional MCP (Model Context Protocol) client module for external tool servers.
252+
- [ ] Add context-window management strategies (trimming and summarization) for long conversations.
251253

252254
## Skills
253255

Sources/SwiftAgentCore/Core/AgentLoop.swift

Lines changed: 183 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -1,5 +1,53 @@
11
import Foundation
22

3+
private enum AgentLoopRuntimeError: LocalizedError {
4+
case llmCallTimedOut(seconds: TimeInterval)
5+
case totalTimeoutExceeded(seconds: TimeInterval)
6+
7+
var errorDescription: String? {
8+
switch self {
9+
case .llmCallTimedOut(let seconds):
10+
return "LLM call timed out after \(String(format: "%.2f", seconds))s."
11+
case .totalTimeoutExceeded(let seconds):
12+
return "Agent loop timed out after \(String(format: "%.2f", seconds))s."
13+
}
14+
}
15+
}
16+
17+
private enum TimeoutKind {
18+
case llmCall
19+
case total
20+
}
21+
22+
private func nanoseconds(from seconds: TimeInterval) -> UInt64 {
23+
guard seconds > 0 else { return 0 }
24+
let value = seconds * 1_000_000_000
25+
if value >= Double(UInt64.max) {
26+
return UInt64.max
27+
}
28+
return UInt64(value.rounded())
29+
}
30+
31+
private func retryDelayIfRetryable(error: Error, fallbackDelay: TimeInterval) -> TimeInterval? {
32+
func decision(statusCode: Int, retryAfter: TimeInterval?) -> TimeInterval? {
33+
guard statusCode == 429 || statusCode >= 500 else { return nil }
34+
return max(0, retryAfter ?? fallbackDelay)
35+
}
36+
37+
switch error {
38+
case OpenAICompatibleProviderError.requestFailed(let statusCode, _, let retryAfter):
39+
return decision(statusCode: statusCode, retryAfter: retryAfter)
40+
case GeminiProviderError.requestFailed(let statusCode, _, let retryAfter):
41+
return decision(statusCode: statusCode, retryAfter: retryAfter)
42+
case MiniMaxAnthropicProviderError.requestFailed(let statusCode, _, let retryAfter):
43+
return decision(statusCode: statusCode, retryAfter: retryAfter)
44+
case OpenAIResponsesProviderError.requestFailed(let statusCode, _):
45+
return decision(statusCode: statusCode, retryAfter: nil)
46+
default:
47+
return nil
48+
}
49+
}
50+
351
public func runAgentLoop(
452
config: AgentLoopConfig,
553
initialMessages: [LLMMessage],
@@ -12,6 +60,138 @@ public func runAgentLoop(
1260
continuation.yield(event)
1361
}
1462

63+
let loopStartedAt = Date()
64+
65+
func remainingTotalTimeout() -> TimeInterval? {
66+
guard let totalTimeout = config.totalTimeout else { return nil }
67+
return totalTimeout - Date().timeIntervalSince(loopStartedAt)
68+
}
69+
70+
func ensureTotalTimeoutNotExceeded() throws {
71+
guard let totalTimeout = config.totalTimeout else { return }
72+
let elapsed = Date().timeIntervalSince(loopStartedAt)
73+
if elapsed >= totalTimeout {
74+
throw AgentLoopRuntimeError.totalTimeoutExceeded(seconds: totalTimeout)
75+
}
76+
}
77+
78+
func effectiveCallTimeout() throws -> (seconds: TimeInterval, kind: TimeoutKind)? {
79+
let llmCallTimeout = config.llmCallTimeout
80+
guard let remaining = remainingTotalTimeout() else {
81+
guard let llmCallTimeout else { return nil }
82+
if llmCallTimeout <= 0 {
83+
throw AgentLoopRuntimeError.llmCallTimedOut(seconds: llmCallTimeout)
84+
}
85+
return (llmCallTimeout, .llmCall)
86+
}
87+
88+
guard let totalTimeout = config.totalTimeout else { return nil }
89+
if remaining <= 0 {
90+
throw AgentLoopRuntimeError.totalTimeoutExceeded(seconds: totalTimeout)
91+
}
92+
93+
guard let llmCallTimeout else {
94+
return (remaining, .total)
95+
}
96+
97+
if llmCallTimeout <= 0 {
98+
throw AgentLoopRuntimeError.llmCallTimedOut(seconds: llmCallTimeout)
99+
}
100+
101+
if llmCallTimeout <= remaining {
102+
return (llmCallTimeout, .llmCall)
103+
}
104+
return (remaining, .total)
105+
}
106+
107+
func sendMessageWithTimeout(
108+
system: String,
109+
messages: [LLMMessage],
110+
tools: [ToolDefinition],
111+
onTextDelta: @escaping (String) -> Void
112+
) async throws -> LLMResponse {
113+
guard let timeout = try effectiveCallTimeout() else {
114+
return try await config.provider.sendMessage(
115+
system: system,
116+
messages: messages,
117+
tools: tools,
118+
onTextDelta: onTextDelta
119+
)
120+
}
121+
122+
return try await withThrowingTaskGroup(of: LLMResponse.self) { group in
123+
group.addTask {
124+
try await config.provider.sendMessage(
125+
system: system,
126+
messages: messages,
127+
tools: tools,
128+
onTextDelta: onTextDelta
129+
)
130+
}
131+
group.addTask {
132+
try await Task.sleep(nanoseconds: nanoseconds(from: timeout.seconds))
133+
switch timeout.kind {
134+
case .llmCall:
135+
throw AgentLoopRuntimeError.llmCallTimedOut(seconds: timeout.seconds)
136+
case .total:
137+
throw AgentLoopRuntimeError.totalTimeoutExceeded(
138+
seconds: config.totalTimeout ?? timeout.seconds
139+
)
140+
}
141+
}
142+
143+
guard let result = try await group.next() else {
144+
throw AgentLoopRuntimeError.llmCallTimedOut(seconds: timeout.seconds)
145+
}
146+
group.cancelAll()
147+
return result
148+
}
149+
}
150+
151+
func sendMessageWithRetry(
152+
system: String,
153+
messages: [LLMMessage],
154+
tools: [ToolDefinition],
155+
onTextDelta: @escaping (String) -> Void
156+
) async throws -> LLMResponse {
157+
var retriesRemaining = config.maxRetries
158+
159+
while true {
160+
try ensureTotalTimeoutNotExceeded()
161+
do {
162+
return try await sendMessageWithTimeout(
163+
system: system,
164+
messages: messages,
165+
tools: tools,
166+
onTextDelta: onTextDelta
167+
)
168+
} catch {
169+
guard !Task.isCancelled else { throw error }
170+
guard retriesRemaining > 0 else { throw error }
171+
guard
172+
let retryDelay = retryDelayIfRetryable(
173+
error: error,
174+
fallbackDelay: config.retryDelay
175+
)
176+
else {
177+
throw error
178+
}
179+
180+
retriesRemaining -= 1
181+
guard retryDelay > 0 else { continue }
182+
183+
if let remaining = remainingTotalTimeout() {
184+
let totalTimeout = config.totalTimeout ?? retryDelay
185+
if remaining <= 0 || remaining < retryDelay {
186+
throw AgentLoopRuntimeError.totalTimeoutExceeded(seconds: totalTimeout)
187+
}
188+
}
189+
190+
try await Task.sleep(nanoseconds: nanoseconds(from: retryDelay))
191+
}
192+
}
193+
}
194+
15195
emit(.agentStart)
16196

17197
let registry = ToolRegistry(tools: config.tools)
@@ -36,8 +216,9 @@ public func runAgentLoop(
36216
emit(.messageStart(role: .assistant))
37217

38218
do {
219+
try ensureTotalTimeoutNotExceeded()
39220
let systemPrompt = await config.buildSystemPrompt()
40-
let response = try await config.provider.sendMessage(
221+
let response = try await sendMessageWithRetry(
41222
system: systemPrompt,
42223
messages: conversation,
43224
tools: registry.definitions,
@@ -66,6 +247,7 @@ public func runAgentLoop(
66247
var results: [ContentBlock] = []
67248

68249
for toolCall in toolUses where !Task.isCancelled {
250+
try ensureTotalTimeoutNotExceeded()
69251
guard let tool = registry.tool(named: toolCall.name) else {
70252
let text = "Tool '\(toolCall.name)' is not registered."
71253
emit(.toolExecutionEnd(name: toolCall.name, result: text, isError: true))

Sources/SwiftAgentCore/Core/AgentTypes.swift

Lines changed: 13 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -7,20 +7,32 @@ public struct AgentLoopConfig {
77
public var confirmationHandler: (String, String) async -> Bool
88
public var getSteeringMessages: () async -> [LLMMessage]
99
public var getFollowUpMessages: () async -> [LLMMessage]
10+
public var maxRetries: Int
11+
public var retryDelay: TimeInterval
12+
public var llmCallTimeout: TimeInterval?
13+
public var totalTimeout: TimeInterval?
1014

1115
public init(
1216
provider: LLMProvider,
1317
tools: [AgentTool],
1418
buildSystemPrompt: @escaping () async -> String,
1519
confirmationHandler: @escaping (String, String) async -> Bool,
1620
getSteeringMessages: @escaping () async -> [LLMMessage] = { [] },
17-
getFollowUpMessages: @escaping () async -> [LLMMessage] = { [] }
21+
getFollowUpMessages: @escaping () async -> [LLMMessage] = { [] },
22+
maxRetries: Int = 2,
23+
retryDelay: TimeInterval = 1.0,
24+
llmCallTimeout: TimeInterval? = nil,
25+
totalTimeout: TimeInterval? = nil
1826
) {
1927
self.provider = provider
2028
self.tools = tools
2129
self.buildSystemPrompt = buildSystemPrompt
2230
self.confirmationHandler = confirmationHandler
2331
self.getSteeringMessages = getSteeringMessages
2432
self.getFollowUpMessages = getFollowUpMessages
33+
self.maxRetries = max(0, maxRetries)
34+
self.retryDelay = max(0, retryDelay)
35+
self.llmCallTimeout = llmCallTimeout
36+
self.totalTimeout = totalTimeout
2537
}
2638
}

0 commit comments

Comments
 (0)