Cap tool-call rounds at 8 to prevent context overflow
ober
039ad5c42051abf3518f3a8b1125ed6fa0338ffd
--- a/lib/jcode/core/agent.sls +++ b/lib/jcode/core/agent.sls @@ -24,6 +24,7 @@ (def current-stream-cb (make-parameter #f)) (def current-tool-cb (make-parameter #f)) (def current-usage-cb (make-parameter #f)) + (def *max-tool-rounds* 8) (def (agent-run session-id user-input) (log-info logger "agent-run" `((session . ,session-id))) (let ([existing (session-get-messages session-id)]) @@ -37,9 +38,13 @@ (if (current-stream-cb) (agent-loop-stream session-id - (session-get-messages session-id)) - (agent-loop session-id (session-get-messages session-id)))) - (def (agent-loop session-id messages) + (session-get-messages session-id) + 0) + (agent-loop + session-id + (session-get-messages session-id) + 0))) + (def (agent-loop session-id messages round) (let* ([provider (get-current-provider)] [tools (get-tool-schemas)] [response (provider-chat provider messages tools)]) @@ -48,15 +53,22 @@ "got-response" `((role . ,(message-role response)))) (session-add-message session-id response) - (if (message-tool-calls response) - (let ([results (execute-tool-calls - (message-tool-calls response))]) - (for-each - (lambda (result) (session-add-message session-id result)) - results) - (agent-loop session-id (session-get-messages session-id))) - response))) - (def (agent-loop-stream session-id messages) + (cond + [(not (message-tool-calls response)) response] + [(>= round *max-tool-rounds*) + (log-warn logger "max-rounds" `((round . ,round))) + response] + [else + (let ([results (execute-tool-calls + (message-tool-calls response))]) + (for-each + (lambda (result) (session-add-message session-id result)) + results) + (agent-loop + session-id + (session-get-messages session-id) + (+ round 1)))]))) + (def (agent-loop-stream session-id messages round) (let* ([provider (get-current-provider)] [tools (get-tool-schemas)]) (let-values ([(content tool-calls usage) @@ -71,15 +83,20 @@ (if (string=? content "") #f content) (if (null? tool-calls) #f tool-calls))]) (session-add-message session-id response) - (if (null? tool-calls) - response - (let ([results (execute-tool-calls tool-calls)]) - (for-each - (lambda (r) (session-add-message session-id r)) - results) - (agent-loop-stream - session-id - (session-get-messages session-id)))))))) + (cond + [(null? tool-calls) response] + [(>= round *max-tool-rounds*) + (log-warn logger "max-rounds" `((round . ,round))) + response] + [else + (let ([results (execute-tool-calls tool-calls)]) + (for-each + (lambda (r) (session-add-message session-id r)) + results) + (agent-loop-stream + session-id + (session-get-messages session-id) + (+ round 1)))]))))) (def (execute-tool-calls tool-calls) (log-info logger @@ -113,34 +130,48 @@ (make-system-message (system-prompt)) (make-user-message user-input))]) (if (current-stream-cb) - (agent-chat-loop-stream provider messages tools) - (agent-chat-loop provider messages tools)))) - (def (agent-chat-loop provider messages tools) + (agent-chat-loop-stream provider messages tools 0) + (agent-chat-loop provider messages tools 0)))) + (def (agent-chat-loop provider messages tools round) (let ([response (provider-chat provider messages tools)]) - (if (message-tool-calls response) - (let* ([results (execute-tool-calls - (message-tool-calls response))] - [new-messages (append - messages - (list response) - results)]) - (agent-chat-loop provider new-messages tools)) - (message-content response)))) - (def (agent-chat-loop-stream provider messages tools) + (cond + [(not (message-tool-calls response)) + (message-content response)] + [(>= round *max-tool-rounds*) + (or (message-content response) "")] + [else + (let* ([results (execute-tool-calls + (message-tool-calls response))] + [new-messages (append + messages + (list response) + results)]) + (agent-chat-loop + provider + new-messages + tools + (+ round 1)))]))) + (def (agent-chat-loop-stream provider messages tools round) (let-values ([(content tool-calls usage) (provider-stream-chat provider messages tools (current-stream-cb))]) - (if (null? tool-calls) - content - (let* ([response (make-assistant-message - (if (string=? content "") #f content) - tool-calls)] - [results (execute-tool-calls tool-calls)] - [new-msgs (append messages (list response) results)]) - (agent-chat-loop-stream provider new-msgs tools))))) + (cond + [(null? tool-calls) content] + [(>= round *max-tool-rounds*) content] + [else + (let* ([response (make-assistant-message + (if (string=? content "") #f content) + tool-calls)] + [results (execute-tool-calls tool-calls)] + [new-msgs (append messages (list response) results)]) + (agent-chat-loop-stream + provider + new-msgs + tools + (+ round 1)))]))) (def (agent-step messages) (let* ([provider (get-current-provider)] [tools (get-tool-schemas)]) --- a/src/jcode/core/agent.ss +++ b/src/jcode/core/agent.ss @@ -46,6 +46,8 @@ Prefer using the edit tool over write for modifying existing files." (current-di (def current-tool-cb (make-parameter #f)) (def current-usage-cb (make-parameter #f)) +(def *max-tool-rounds* 8) + (def (agent-run session-id user-input) (log-info logger "agent-run" `((session . ,session-id))) (let ((existing (session-get-messages session-id))) @@ -53,24 +55,28 @@ Prefer using the edit tool over write for modifying existing files." (current-di (session-add-message session-id (make-system-message (system-prompt))))) (session-add-message session-id (make-user-message user-input)) (if (current-stream-cb) - (agent-loop-stream session-id (session-get-messages session-id)) - (agent-loop session-id (session-get-messages session-id)))) + (agent-loop-stream session-id (session-get-messages session-id) 0) + (agent-loop session-id (session-get-messages session-id) 0))) -(def (agent-loop session-id messages) +(def (agent-loop session-id messages round) (let* ((provider (get-current-provider)) (tools (get-tool-schemas)) (response (provider-chat provider messages tools))) (log-debug logger "got-response" `((role . ,(message-role response)))) (session-add-message session-id response) - (if (message-tool-calls response) - (let ((results (execute-tool-calls (message-tool-calls response)))) - (for-each - (lambda (result) (session-add-message session-id result)) - results) - (agent-loop session-id (session-get-messages session-id))) - response))) - -(def (agent-loop-stream session-id messages) + (cond + ((not (message-tool-calls response)) response) + ((>= round *max-tool-rounds*) + (log-warn logger "max-rounds" `((round . ,round))) + response) + (else + (let ((results (execute-tool-calls (message-tool-calls response)))) + (for-each + (lambda (result) (session-add-message session-id result)) + results) + (agent-loop session-id (session-get-messages session-id) (+ round 1))))))) + +(def (agent-loop-stream session-id messages round) ;; Streaming version: calls (current-stream-cb) for each text token. (let* ((provider (get-current-provider)) (tools (get-tool-schemas))) @@ -82,11 +88,15 @@ Prefer using the edit tool over write for modifying existing files." (current-di (if (string=? content "") #f content) (if (null? tool-calls) #f tool-calls)))) (session-add-message session-id response) - (if (null? tool-calls) - response - (let ((results (execute-tool-calls tool-calls))) - (for-each (lambda (r) (session-add-message session-id r)) results) - (agent-loop-stream session-id (session-get-messages session-id)))))))) + (cond + ((null? tool-calls) response) + ((>= round *max-tool-rounds*) + (log-warn logger "max-rounds" `((round . ,round))) + response) + (else + (let ((results (execute-tool-calls tool-calls))) + (for-each (lambda (r) (session-add-message session-id r)) results) + (agent-loop-stream session-id (session-get-messages session-id) (+ round 1))))))))) (def (execute-tool-calls tool-calls) (log-info logger "executing-tools" `((count . ,(length tool-calls)))) @@ -117,28 +127,32 @@ Prefer using the edit tool over write for modifying existing files." (current-di (make-system-message (system-prompt)) (make-user-message user-input)))) (if (current-stream-cb) - (agent-chat-loop-stream provider messages tools) - (agent-chat-loop provider messages tools)))) + (agent-chat-loop-stream provider messages tools 0) + (agent-chat-loop provider messages tools 0)))) -(def (agent-chat-loop provider messages tools) +(def (agent-chat-loop provider messages tools round) (let ((response (provider-chat provider messages tools))) - (if (message-tool-calls response) - (let* ((results (execute-tool-calls (message-tool-calls response))) - (new-messages (append messages (list response) results))) - (agent-chat-loop provider new-messages tools)) - (message-content response)))) - -(def (agent-chat-loop-stream provider messages tools) + (cond + ((not (message-tool-calls response)) (message-content response)) + ((>= round *max-tool-rounds*) (or (message-content response) "")) + (else + (let* ((results (execute-tool-calls (message-tool-calls response))) + (new-messages (append messages (list response) results))) + (agent-chat-loop provider new-messages tools (+ round 1))))))) + +(def (agent-chat-loop-stream provider messages tools round) (let-values (((content tool-calls usage) (provider-stream-chat provider messages tools (current-stream-cb)))) - (if (null? tool-calls) - content - (let* ((response (make-assistant-message - (if (string=? content "") #f content) - tool-calls)) - (results (execute-tool-calls tool-calls)) - (new-msgs (append messages (list response) results))) - (agent-chat-loop-stream provider new-msgs tools))))) + (cond + ((null? tool-calls) content) + ((>= round *max-tool-rounds*) content) + (else + (let* ((response (make-assistant-message + (if (string=? content "") #f content) + tool-calls)) + (results (execute-tool-calls tool-calls)) + (new-msgs (append messages (list response) results))) + (agent-chat-loop-stream provider new-msgs tools (+ round 1))))))) (def (agent-step messages) (let* ((provider (get-current-provider))