fix: prevent Ollama provider from hanging on tool-calling requests (#7723)
Signed-off-by: fre <anonwurcod@proton.me> Signed-off-by: fresh3nough <nicholasanthony742@gmail.com> Signed-off-by: Douwe Osinga <douwe@squareup.com> Co-authored-by: Douwe Osinga <douwe@squareup.com>
This commit is contained in:
Generated
+13
-10
@@ -4342,6 +4342,7 @@ dependencies = [
|
|||||||
"base64 0.22.1",
|
"base64 0.22.1",
|
||||||
"blake3",
|
"blake3",
|
||||||
"byteorder",
|
"byteorder",
|
||||||
|
"bytes",
|
||||||
"candle-core",
|
"candle-core",
|
||||||
"candle-nn",
|
"candle-nn",
|
||||||
"candle-transformers",
|
"candle-transformers",
|
||||||
@@ -4358,6 +4359,7 @@ dependencies = [
|
|||||||
"fs2",
|
"fs2",
|
||||||
"futures",
|
"futures",
|
||||||
"goose-test-support",
|
"goose-test-support",
|
||||||
|
"http 1.4.0",
|
||||||
"ignore",
|
"ignore",
|
||||||
"include_dir",
|
"include_dir",
|
||||||
"indexmap 2.13.0",
|
"indexmap 2.13.0",
|
||||||
@@ -5787,9 +5789,9 @@ checksum = "b6d2cec3eae94f9f509c767b45932f1ada8350c4bdb85af2fcab4a3c14807981"
|
|||||||
|
|
||||||
[[package]]
|
[[package]]
|
||||||
name = "libredox"
|
name = "libredox"
|
||||||
version = "0.1.14"
|
version = "0.1.15"
|
||||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||||
checksum = "1744e39d1d6a9948f4f388969627434e31128196de472883b39f148769bfe30a"
|
checksum = "7ddbf48fd451246b1f8c2610bd3b4ac0cc6e149d89832867093ab69a17194f08"
|
||||||
dependencies = [
|
dependencies = [
|
||||||
"bitflags 2.11.0",
|
"bitflags 2.11.0",
|
||||||
"libc",
|
"libc",
|
||||||
@@ -6344,9 +6346,9 @@ dependencies = [
|
|||||||
|
|
||||||
[[package]]
|
[[package]]
|
||||||
name = "num-conv"
|
name = "num-conv"
|
||||||
version = "0.2.0"
|
version = "0.2.1"
|
||||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||||
checksum = "cf97ec579c3c42f953ef76dbf8d55ac91fb219dde70e49aa4a6b7d74e9919050"
|
checksum = "c6673768db2d862beb9b39a78fdcb1a69439615d5794a1be50caa9bc92c81967"
|
||||||
|
|
||||||
[[package]]
|
[[package]]
|
||||||
name = "num-derive"
|
name = "num-derive"
|
||||||
@@ -6960,9 +6962,9 @@ dependencies = [
|
|||||||
|
|
||||||
[[package]]
|
[[package]]
|
||||||
name = "pctx_executor"
|
name = "pctx_executor"
|
||||||
version = "0.2.0"
|
version = "0.2.1"
|
||||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||||
checksum = "52c6f6e9ae409b20157d7f60623e32de6627b7a80a159e3ee5c1472c4f36ba64"
|
checksum = "459b68c42775a94e0e160537e644f9a8d4ab7041fa724976b08382eb386836e1"
|
||||||
dependencies = [
|
dependencies = [
|
||||||
"deno_core",
|
"deno_core",
|
||||||
"deno_resolver",
|
"deno_resolver",
|
||||||
@@ -6984,13 +6986,14 @@ dependencies = [
|
|||||||
|
|
||||||
[[package]]
|
[[package]]
|
||||||
name = "pctx_registry"
|
name = "pctx_registry"
|
||||||
version = "0.1.0"
|
version = "0.1.1"
|
||||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||||
checksum = "9a49c948ffc8c07357e76b2e0008503c9fafaf49d91fd7e00e672bbd7aabd157"
|
checksum = "2f5a53bf73ba98a5352e788391f47c33af096c696c64e130110b85cc51336e20"
|
||||||
dependencies = [
|
dependencies = [
|
||||||
"deno_error",
|
"deno_error",
|
||||||
"pctx_config",
|
"pctx_config",
|
||||||
"rmcp",
|
"rmcp",
|
||||||
|
"serde",
|
||||||
"serde_json",
|
"serde_json",
|
||||||
"thiserror 2.0.18",
|
"thiserror 2.0.18",
|
||||||
"tokio",
|
"tokio",
|
||||||
@@ -11297,9 +11300,9 @@ checksum = "7df058c713841ad818f1dc5d3fd88063241cc61f49f5fbea4b951e8cf5a8d71d"
|
|||||||
|
|
||||||
[[package]]
|
[[package]]
|
||||||
name = "unicode-segmentation"
|
name = "unicode-segmentation"
|
||||||
version = "1.12.0"
|
version = "1.13.1"
|
||||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||||
checksum = "f6ccf251212114b54433ec949fd6a7841275f9ada20dddd2f29e9ceea4501493"
|
checksum = "da36089a805484bcccfffe0739803392c8298778a2d2f09febf76fac5ad9025b"
|
||||||
|
|
||||||
[[package]]
|
[[package]]
|
||||||
name = "unicode-width"
|
name = "unicode-width"
|
||||||
|
|||||||
@@ -215,6 +215,8 @@ env-lock = { workspace = true }
|
|||||||
rmcp = { workspace = true, features = ["transport-streamable-http-server"] }
|
rmcp = { workspace = true, features = ["transport-streamable-http-server"] }
|
||||||
opentelemetry_sdk = { workspace = true, features = ["testing"] }
|
opentelemetry_sdk = { workspace = true, features = ["testing"] }
|
||||||
goose-test-support = { path = "../goose-test-support" }
|
goose-test-support = { path = "../goose-test-support" }
|
||||||
|
bytes.workspace = true
|
||||||
|
http.workspace = true
|
||||||
|
|
||||||
[[example]]
|
[[example]]
|
||||||
name = "agent"
|
name = "agent"
|
||||||
|
|||||||
@@ -68,14 +68,28 @@ fn resolve_ollama_num_ctx(model_config: &ModelConfig) -> Option<usize> {
|
|||||||
}
|
}
|
||||||
|
|
||||||
fn apply_ollama_options(payload: &mut Value, model_config: &ModelConfig) {
|
fn apply_ollama_options(payload: &mut Value, model_config: &ModelConfig) {
|
||||||
let Some(limit) = resolve_ollama_num_ctx(model_config) else {
|
|
||||||
return;
|
|
||||||
};
|
|
||||||
|
|
||||||
if let Some(obj) = payload.as_object_mut() {
|
if let Some(obj) = payload.as_object_mut() {
|
||||||
let options = obj.entry("options").or_insert_with(|| json!({}));
|
// Ollama does not support stream_options; remove it to prevent hangs.
|
||||||
if let Some(options_obj) = options.as_object_mut() {
|
obj.remove("stream_options");
|
||||||
options_obj.insert("num_ctx".to_string(), json!(limit));
|
|
||||||
|
// Convert max_completion_tokens / max_tokens to Ollama's options.num_predict.
|
||||||
|
// Reasoning models emit max_completion_tokens; non-reasoning models emit max_tokens.
|
||||||
|
let max_tokens = obj
|
||||||
|
.remove("max_completion_tokens")
|
||||||
|
.or_else(|| obj.remove("max_tokens"));
|
||||||
|
if let Some(max_tokens) = max_tokens {
|
||||||
|
let options = obj.entry("options").or_insert_with(|| json!({}));
|
||||||
|
if let Some(options_obj) = options.as_object_mut() {
|
||||||
|
options_obj.entry("num_predict").or_insert(max_tokens);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Apply num_ctx from context limit settings.
|
||||||
|
if let Some(limit) = resolve_ollama_num_ctx(model_config) {
|
||||||
|
let options = obj.entry("options").or_insert_with(|| json!({}));
|
||||||
|
if let Some(options_obj) = options.as_object_mut() {
|
||||||
|
options_obj.insert("num_ctx".to_string(), json!(limit));
|
||||||
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -300,9 +314,49 @@ impl Provider for OllamaProvider {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
/// Per-chunk timeout for Ollama streaming responses.
|
||||||
|
/// If no new raw SSE data arrives within this duration, the connection is considered dead.
|
||||||
|
const OLLAMA_CHUNK_TIMEOUT_SECS: u64 = 30;
|
||||||
|
|
||||||
|
/// Wraps a line stream with a per-item timeout at the raw SSE level.
|
||||||
|
/// This detects dead connections without false-positive stalls during long
|
||||||
|
/// tool-call generations where response_to_streaming_message_ollama buffers.
|
||||||
|
fn with_line_timeout(
|
||||||
|
stream: impl futures::Stream<Item = anyhow::Result<String>> + Unpin + Send + 'static,
|
||||||
|
timeout_secs: u64,
|
||||||
|
) -> std::pin::Pin<Box<dyn futures::Stream<Item = anyhow::Result<String>> + Send>> {
|
||||||
|
let timeout = Duration::from_secs(timeout_secs);
|
||||||
|
Box::pin(try_stream! {
|
||||||
|
let mut stream = stream;
|
||||||
|
|
||||||
|
// Allow time-to-first-token to be governed by the request timeout.
|
||||||
|
// Only enforce per-chunk timeout after first SSE line arrives.
|
||||||
|
match stream.next().await {
|
||||||
|
Some(first_item) => yield first_item?,
|
||||||
|
None => return,
|
||||||
|
}
|
||||||
|
loop {
|
||||||
|
match tokio::time::timeout(timeout, stream.next()).await {
|
||||||
|
Ok(Some(item)) => yield item?,
|
||||||
|
Ok(None) => break,
|
||||||
|
Err(_) => {
|
||||||
|
Err::<(), anyhow::Error>(anyhow::anyhow!(
|
||||||
|
"Ollama stream stalled: no data received for {}s. \
|
||||||
|
This may indicate the model is overwhelmed by the request payload. \
|
||||||
|
Try a smaller model or reduce the number of tools.",
|
||||||
|
timeout_secs
|
||||||
|
))?;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
/// Ollama-specific streaming handler with XML tool call fallback.
|
/// Ollama-specific streaming handler with XML tool call fallback.
|
||||||
/// Uses the Ollama format module which buffers text when XML tool calls are detected,
|
/// Uses the Ollama format module which buffers text when XML tool calls are detected,
|
||||||
/// preventing duplicate content from being emitted to the UI.
|
/// preventing duplicate content from being emitted to the UI.
|
||||||
|
/// Timeout is applied at the raw SSE line level via with_line_timeout so that
|
||||||
|
/// buffering inside response_to_streaming_message_ollama does not cause false stalls.
|
||||||
fn stream_ollama(response: Response, mut log: RequestLog) -> Result<MessageStream, ProviderError> {
|
fn stream_ollama(response: Response, mut log: RequestLog) -> Result<MessageStream, ProviderError> {
|
||||||
let stream = response.bytes_stream().map_err(std::io::Error::other);
|
let stream = response.bytes_stream().map_err(std::io::Error::other);
|
||||||
|
|
||||||
@@ -311,8 +365,10 @@ fn stream_ollama(response: Response, mut log: RequestLog) -> Result<MessageStrea
|
|||||||
let framed = FramedRead::new(stream_reader, LinesCodec::new())
|
let framed = FramedRead::new(stream_reader, LinesCodec::new())
|
||||||
.map_err(Error::from);
|
.map_err(Error::from);
|
||||||
|
|
||||||
let message_stream = response_to_streaming_message_ollama(framed);
|
let timed_lines = with_line_timeout(framed, OLLAMA_CHUNK_TIMEOUT_SECS);
|
||||||
|
let message_stream = response_to_streaming_message_ollama(timed_lines);
|
||||||
pin!(message_stream);
|
pin!(message_stream);
|
||||||
|
|
||||||
while let Some(message) = message_stream.next().await {
|
while let Some(message) = message_stream.next().await {
|
||||||
let (message, usage) = message.map_err(|e|
|
let (message, usage) = message.map_err(|e|
|
||||||
ProviderError::RequestFailed(format!("Stream decode error: {}", e))
|
ProviderError::RequestFailed(format!("Stream decode error: {}", e))
|
||||||
@@ -359,6 +415,131 @@ mod tests {
|
|||||||
assert!(payload.get("options").is_none());
|
assert!(payload.get("options").is_none());
|
||||||
}
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn test_raw_create_request_contains_unsupported_ollama_fields() {
|
||||||
|
use crate::providers::formats::ollama::create_request;
|
||||||
|
use crate::providers::utils::ImageFormat;
|
||||||
|
|
||||||
|
let model_config = ModelConfig::new("llama3.1")
|
||||||
|
.unwrap()
|
||||||
|
.with_max_tokens(Some(4096));
|
||||||
|
let messages = vec![crate::conversation::message::Message::user().with_text("hi")];
|
||||||
|
|
||||||
|
let payload = create_request(
|
||||||
|
&model_config,
|
||||||
|
"You are a helpful assistant.",
|
||||||
|
&messages,
|
||||||
|
&[],
|
||||||
|
&ImageFormat::OpenAi,
|
||||||
|
true,
|
||||||
|
)
|
||||||
|
.unwrap();
|
||||||
|
|
||||||
|
assert!(
|
||||||
|
payload.get("stream_options").is_some(),
|
||||||
|
"create_request should produce stream_options (unsupported by Ollama)"
|
||||||
|
);
|
||||||
|
assert!(
|
||||||
|
payload.get("max_tokens").is_some(),
|
||||||
|
"create_request should produce max_tokens (unsupported by Ollama)"
|
||||||
|
);
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn test_apply_ollama_options_strips_unsupported_fields() {
|
||||||
|
use crate::providers::formats::ollama::create_request;
|
||||||
|
use crate::providers::utils::ImageFormat;
|
||||||
|
|
||||||
|
let _guard = env_lock::lock_env([("GOOSE_INPUT_LIMIT", None::<&str>)]);
|
||||||
|
let model_config = ModelConfig::new("llama3.1")
|
||||||
|
.unwrap()
|
||||||
|
.with_max_tokens(Some(4096));
|
||||||
|
let messages = vec![crate::conversation::message::Message::user().with_text("hi")];
|
||||||
|
|
||||||
|
let mut payload = create_request(
|
||||||
|
&model_config,
|
||||||
|
"You are a helpful assistant.",
|
||||||
|
&messages,
|
||||||
|
&[],
|
||||||
|
&ImageFormat::OpenAi,
|
||||||
|
true,
|
||||||
|
)
|
||||||
|
.unwrap();
|
||||||
|
|
||||||
|
apply_ollama_options(&mut payload, &model_config);
|
||||||
|
|
||||||
|
assert!(
|
||||||
|
payload.get("stream_options").is_none(),
|
||||||
|
"stream_options should be removed for Ollama"
|
||||||
|
);
|
||||||
|
assert!(
|
||||||
|
payload.get("max_tokens").is_none(),
|
||||||
|
"max_tokens should be removed for Ollama"
|
||||||
|
);
|
||||||
|
assert!(
|
||||||
|
payload.get("max_completion_tokens").is_none(),
|
||||||
|
"max_completion_tokens should be removed for Ollama"
|
||||||
|
);
|
||||||
|
assert_eq!(
|
||||||
|
payload["options"]["num_predict"], 4096,
|
||||||
|
"max_tokens should be moved to options.num_predict"
|
||||||
|
);
|
||||||
|
assert_eq!(payload["stream"], true, "stream field should be preserved");
|
||||||
|
}
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
async fn test_stream_ollama_timeout_on_stall() {
|
||||||
|
use std::convert::Infallible;
|
||||||
|
|
||||||
|
let (tx, rx) = tokio::sync::mpsc::channel::<Result<bytes::Bytes, Infallible>>(1);
|
||||||
|
tx.send(Ok(bytes::Bytes::from(
|
||||||
|
"data: {\"choices\":[{\"delta\":{\"content\":\"hi\"},\"index\":0}],\
|
||||||
|
\"model\":\"test\",\"object\":\"chat.completion.chunk\",\"created\":0}\n",
|
||||||
|
)))
|
||||||
|
.await
|
||||||
|
.unwrap();
|
||||||
|
let stream = tokio_stream::wrappers::ReceiverStream::new(rx);
|
||||||
|
let body = reqwest::Body::wrap_stream(stream);
|
||||||
|
let response = http::Response::builder().status(200).body(body).unwrap();
|
||||||
|
let response: reqwest::Response = response.into();
|
||||||
|
|
||||||
|
let log = RequestLog::start(
|
||||||
|
&ModelConfig::new("test").unwrap(),
|
||||||
|
&json!({"model": "test"}),
|
||||||
|
)
|
||||||
|
.unwrap();
|
||||||
|
|
||||||
|
let mut msg_stream = stream_ollama(response, log).unwrap();
|
||||||
|
|
||||||
|
let result =
|
||||||
|
tokio::time::timeout(Duration::from_secs(OLLAMA_CHUNK_TIMEOUT_SECS + 5), async {
|
||||||
|
let mut last_err = None;
|
||||||
|
while let Some(item) = msg_stream.next().await {
|
||||||
|
if let Err(e) = item {
|
||||||
|
last_err = Some(e);
|
||||||
|
break;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
last_err
|
||||||
|
})
|
||||||
|
.await;
|
||||||
|
|
||||||
|
match result {
|
||||||
|
Ok(Some(err)) => {
|
||||||
|
let err_msg = err.to_string();
|
||||||
|
assert!(
|
||||||
|
err_msg.contains("stream stalled"),
|
||||||
|
"Expected stall timeout error, got: {}",
|
||||||
|
err_msg
|
||||||
|
);
|
||||||
|
}
|
||||||
|
Ok(None) => panic!("Expected timeout error but stream completed normally"),
|
||||||
|
Err(_) => panic!("Outer timeout elapsed -- per-chunk timeout did not fire"),
|
||||||
|
}
|
||||||
|
|
||||||
|
drop(tx);
|
||||||
|
}
|
||||||
|
|
||||||
#[test]
|
#[test]
|
||||||
fn test_ollama_retry_config_is_transient_only() {
|
fn test_ollama_retry_config_is_transient_only() {
|
||||||
let config = RetryConfig::new(
|
let config = RetryConfig::new(
|
||||||
|
|||||||
Reference in New Issue
Block a user