diff --git a/src/core/agent.c b/src/core/agent.c index 892f935..28edee3 100644 --- a/src/core/agent.c +++ b/src/core/agent.c @@ -333,12 +333,39 @@ typedef struct agent_run_ctx { char *tool_result_bufs; provider_message_t *messages; size_t total_msgs; + size_t base_msg_count; int history_count; int max_iter; int max_ctx; int ret; } agent_run_ctx_t; +/** True when message content lives in static agent_run buffers (not heap-owned). */ +static int agent_content_is_static(const agent_run_ctx_t *ctx, const char *content) +{ + if (!content) return 1; + if (content == ctx->system_buf) return 1; + if (content == ctx->user_message) return 1; + if (ctx->history_content && content >= ctx->history_content && + content < ctx->history_content + HISTORY_CONTENT_MAX) + return 1; + return 0; +} + +/** Free heap-owned ReAct message slots (assistant text, tool results, tool_calls). */ +static void agent_free_heap_messages(agent_run_ctx_t *ctx, size_t from_idx) +{ + size_t i; + if (!ctx->messages) return; + for (i = from_idx; i < ctx->total_msgs; i++) { + if (!agent_content_is_static(ctx, ctx->messages[i].content)) + free((void *)ctx->messages[i].content); + free_tool_calls_copy(ctx->messages[i].tool_calls, ctx->messages[i].tool_calls_count); + ctx->messages[i].tool_calls = NULL; + ctx->messages[i].tool_calls_count = 0; + } +} + static void agent_oom_msg(agent_run_ctx_t *ctx) { if (ctx->response_buf && ctx->response_size > 0) { @@ -438,32 +465,21 @@ static int agent_react_loop(agent_run_ctx_t *ctx) { int iteration = 0; provider_response_t response = {0}; - char *prev_assistant = NULL; - provider_tool_call_t *prev_calls = NULL; - size_t prev_n = 0; + ctx->base_msg_count = ctx->total_msgs; for (;;) { int err = ctx->provider->chat(ctx->messages, ctx->total_msgs, ctx->tool_defs, ctx->tool_count, &response); if (err != 0) { copy_response_to_buf(response.content, ctx->response_buf, ctx->response_size); provider_response_clear(&response); - free(prev_assistant); - free_tool_calls_copy(prev_calls, prev_n); return -1; } if (response.tool_calls_count == 0 || iteration >= ctx->max_iter) { copy_response_to_buf(response.content, ctx->response_buf, ctx->response_size); provider_response_clear(&response); agent_persist_session(ctx, ctx->response_buf); - free(prev_assistant); - free_tool_calls_copy(prev_calls, prev_n); return 0; } - free(prev_assistant); - free_tool_calls_copy(prev_calls, prev_n); - prev_assistant = NULL; - prev_calls = NULL; - prev_n = 0; { size_t nc = response.tool_calls_count; char *assistant_content; @@ -513,18 +529,25 @@ static int agent_react_loop(agent_run_ctx_t *ctx) new_messages[ctx->total_msgs].tool_calls = our_calls; new_messages[ctx->total_msgs].tool_calls_count = nc; for (size_t k = 0; k < nc; k++) { + char *one_buf = ctx->tool_result_bufs + k * TOOL_RESULT_SIZE; + char *tool_content = strdup(one_buf); + if (!tool_content) { + for (size_t j = 0; j < k; j++) + free((void *)new_messages[ctx->total_msgs + 1 + j].content); + free(new_messages); + free_tool_calls_copy(our_calls, nc); + free(assistant_content); + agent_oom_msg(ctx); + return -1; + } new_messages[ctx->total_msgs + 1 + k].role = "user"; - new_messages[ctx->total_msgs + 1 + k].content = - ctx->tool_result_bufs + k * TOOL_RESULT_SIZE; + new_messages[ctx->total_msgs + 1 + k].content = tool_content; new_messages[ctx->total_msgs + 1 + k].tool_use_id = our_calls[k].id; } free(ctx->messages); ctx->messages = new_messages; ctx->total_msgs = new_count; iteration++; - prev_assistant = assistant_content; - prev_calls = our_calls; - prev_n = nc; } } } @@ -540,6 +563,8 @@ static void agent_run_cleanup(agent_run_ctx_t *ctx) free(ctx->history_msgs); free(ctx->history_roles_buf); free(ctx->tool_result_bufs); + if (ctx->messages && ctx->base_msg_count < ctx->total_msgs) + agent_free_heap_messages(ctx, ctx->base_msg_count); free(ctx->messages); } diff --git a/tests/test_agent.c b/tests/test_agent.c index 1ee8603..504643a 100644 --- a/tests/test_agent.c +++ b/tests/test_agent.c @@ -567,6 +567,102 @@ static int test_agent_unknown_tool_continues(void) return 0; } +static int seq_tool_exec_count; +static int seq_tool_execute(const char *args_json, char *result_buf, size_t max_len) +{ + (void)args_json; + seq_tool_exec_count++; + if (max_len > 0) { + snprintf(result_buf, max_len, "tool_output_%d", seq_tool_exec_count); + result_buf[max_len - 1] = '\0'; + } + return 0; +} +static const agent_tool_t seq_echo_tool = { + .name = "echo", + .description = "Echo test with sequence counter", + .parameters_json = "{}", + .execute = seq_tool_execute, +}; + +static int multi_tool_round_call_count; +static int multi_tool_round_init(const config_t *cfg) { (void)cfg; multi_tool_round_call_count = 0; seq_tool_exec_count = 0; return 0; } +static int multi_tool_round_chat(const provider_message_t *messages, size_t message_count, + const provider_tool_def_t *tools, size_t tool_count, provider_response_t *response) +{ + (void)tools; + (void)tool_count; + size_t i; + int saw_first_tool_output = 0; + response->error = 0; + response->tool_calls = NULL; + response->tool_calls_count = 0; + response->content = NULL; + multi_tool_round_call_count++; + if (multi_tool_round_call_count >= 2) { + for (i = 0; i < message_count; i++) { + if (messages[i].content && strstr(messages[i].content, "tool_output_1") != NULL) { + saw_first_tool_output = 1; + break; + } + } + if (!saw_first_tool_output) { + response->content = strdup("CORRUPTED_TOOL_HISTORY"); + return 0; + } + } + if (multi_tool_round_call_count <= 2) { + response->tool_calls = malloc(sizeof(provider_tool_call_t)); + if (!response->tool_calls) { + response->error = 1; + return -1; + } + response->tool_calls[0].id = strdup("mt1"); + response->tool_calls[0].name = strdup("echo"); + response->tool_calls[0].arguments = strdup("{}"); + response->tool_calls_count = 1; + response->content = strdup(""); + return 0; + } + response->content = strdup("multi tool done"); + return 0; +} +static void multi_tool_round_cleanup(void) {} +static const provider_t multi_tool_round_provider = { + .name = "multi_tool_round", + .init = multi_tool_round_init, + .chat = multi_tool_round_chat, + .cleanup = multi_tool_round_cleanup, +}; + +static int test_react_loop_preserves_prior_tool_results(void) +{ + int failed = 1; + const char *path = "build/test_agent_multi_tool.toml"; + FILE *f = fopen(path, "w"); + ASSERT(f); + fprintf(f, "[agent]\nmodel = \"test\"\nmax_tool_iterations = 5\n"); + fclose(f); + config_t *cfg = NULL; + char errbuf[256]; + if (config_load(path, &cfg, errbuf, sizeof(errbuf)) != 0) goto cleanup; + if (cfg == NULL) goto cleanup; + char response_buf[4096]; + response_buf[0] = '\0'; + int ret = agent_run(cfg, "cli:multitool", "hi", &multi_tool_round_provider, &seq_echo_tool, 1, + response_buf, sizeof(response_buf)); + if (ret != 0) goto cleanup; + if (strstr(response_buf, "multi tool done") == NULL) goto cleanup; + if (strstr(response_buf, "CORRUPTED_TOOL_HISTORY") != NULL) goto cleanup; + if (multi_tool_round_call_count != 3) goto cleanup; + if (seq_tool_exec_count != 2) goto cleanup; + failed = 0; +cleanup: + config_free(cfg); + remove(path); + return failed; +} + static int test_local_offline_note_skipped_for_non_local(void) { const char *path = "build/test_agent_nonlocal_note.toml"; @@ -597,6 +693,7 @@ int main(void) RUN(test_context_assembly_system_prompt_history_memories()); RUN(test_react_loop_tool_then_text()); RUN(test_react_loop_max_iterations()); + RUN(test_react_loop_preserves_prior_tool_results()); RUN(test_session_persisted_after_exchange()); RUN(test_context_compaction_when_history_exceeds_max()); RUN(test_local_offline_note_when_active_is_local());