fix: enable token usage tracking and configurable stream timeout for Ollama provider (#8493)
This commit is contained in:
@@ -603,14 +603,21 @@ pub fn get_usage(usage: &Value) -> Usage {
|
|||||||
.filter(|nested| nested.is_object())
|
.filter(|nested| nested.is_object())
|
||||||
.unwrap_or(usage);
|
.unwrap_or(usage);
|
||||||
|
|
||||||
|
// Try standard OpenAI fields first, then fall back to Ollama-native fields
|
||||||
|
// (prompt_eval_count / eval_count) for compatibility with older Ollama builds
|
||||||
|
// that don't translate to OpenAI field names.
|
||||||
|
// Parse the value before falling back so that present-but-null keys
|
||||||
|
// (e.g. "completion_tokens": null) don't block the fallback.
|
||||||
let input_tokens = usage
|
let input_tokens = usage
|
||||||
.get("prompt_tokens")
|
.get("prompt_tokens")
|
||||||
.and_then(|v| v.as_i64())
|
.and_then(|v| v.as_i64())
|
||||||
|
.or_else(|| usage.get("prompt_eval_count").and_then(|v| v.as_i64()))
|
||||||
.map(|v| v as i32);
|
.map(|v| v as i32);
|
||||||
|
|
||||||
let output_tokens = usage
|
let output_tokens = usage
|
||||||
.get("completion_tokens")
|
.get("completion_tokens")
|
||||||
.and_then(|v| v.as_i64())
|
.and_then(|v| v.as_i64())
|
||||||
|
.or_else(|| usage.get("eval_count").and_then(|v| v.as_i64()))
|
||||||
.map(|v| v as i32);
|
.map(|v| v as i32);
|
||||||
|
|
||||||
let cache_read_input_tokens = usage
|
let cache_read_input_tokens = usage
|
||||||
@@ -2288,4 +2295,47 @@ data: [DONE]"#;
|
|||||||
assert_eq!(messages[2]["content"], "what happened?");
|
assert_eq!(messages[2]["content"], "what happened?");
|
||||||
assert_eq!(messages[3]["tool_calls"].as_array().unwrap().len(), 1);
|
assert_eq!(messages[3]["tool_calls"].as_array().unwrap().len(), 1);
|
||||||
}
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn test_get_usage_with_ollama_native_fields() {
|
||||||
|
// Ollama-native fields should be picked up as fallback
|
||||||
|
let usage = json!({
|
||||||
|
"prompt_eval_count": 42,
|
||||||
|
"eval_count": 128
|
||||||
|
});
|
||||||
|
let result = get_usage(&usage);
|
||||||
|
assert_eq!(result.input_tokens, Some(42));
|
||||||
|
assert_eq!(result.output_tokens, Some(128));
|
||||||
|
assert_eq!(result.total_tokens, Some(170));
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn test_get_usage_prefers_openai_fields_over_ollama() {
|
||||||
|
// Standard OpenAI fields should take precedence
|
||||||
|
let usage = json!({
|
||||||
|
"prompt_tokens": 10,
|
||||||
|
"completion_tokens": 20,
|
||||||
|
"prompt_eval_count": 42,
|
||||||
|
"eval_count": 128
|
||||||
|
});
|
||||||
|
let result = get_usage(&usage);
|
||||||
|
assert_eq!(result.input_tokens, Some(10));
|
||||||
|
assert_eq!(result.output_tokens, Some(20));
|
||||||
|
assert_eq!(result.total_tokens, Some(30));
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn test_get_usage_falls_back_when_openai_fields_are_null() {
|
||||||
|
// When OpenAI fields exist but are null, should fall back to Ollama-native
|
||||||
|
let usage = json!({
|
||||||
|
"prompt_tokens": null,
|
||||||
|
"completion_tokens": null,
|
||||||
|
"prompt_eval_count": 42,
|
||||||
|
"eval_count": 128
|
||||||
|
});
|
||||||
|
let result = get_usage(&usage);
|
||||||
|
assert_eq!(result.input_tokens, Some(42));
|
||||||
|
assert_eq!(result.output_tokens, Some(128));
|
||||||
|
assert_eq!(result.total_tokens, Some(170));
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -68,10 +68,36 @@ fn resolve_ollama_num_ctx(model_config: &ModelConfig) -> Option<usize> {
|
|||||||
input_limit.or(model_config.context_limit)
|
input_limit.or(model_config.context_limit)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
fn resolve_ollama_stream_usage() -> bool {
|
||||||
|
let config = crate::config::Config::global();
|
||||||
|
match config.get_param::<bool>("OLLAMA_STREAM_USAGE") {
|
||||||
|
Ok(val) => val,
|
||||||
|
// Key not set: default to true. Ollama supports stream_options since
|
||||||
|
// mid-2025 and most installs benefit from token usage tracking.
|
||||||
|
Err(crate::config::ConfigError::NotFound(_)) => true,
|
||||||
|
// Invalid value (e.g. "0", "yes", typo): warn and disable stream_options
|
||||||
|
// so users who intended to opt out aren't silently left hanging.
|
||||||
|
Err(e) => {
|
||||||
|
tracing::warn!(
|
||||||
|
"Invalid OLLAMA_STREAM_USAGE value ({}); disabling stream_options. \
|
||||||
|
Use true or false.",
|
||||||
|
e
|
||||||
|
);
|
||||||
|
false
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
fn apply_ollama_options(payload: &mut Value, model_config: &ModelConfig) {
|
fn apply_ollama_options(payload: &mut Value, model_config: &ModelConfig) {
|
||||||
if let Some(obj) = payload.as_object_mut() {
|
if let Some(obj) = payload.as_object_mut() {
|
||||||
// Ollama does not support stream_options; remove it to prevent hangs.
|
// Gate stream_options behind OLLAMA_STREAM_USAGE (default: true).
|
||||||
obj.remove("stream_options");
|
// Older Ollama builds that don't support stream_options may stall before
|
||||||
|
// emitting any SSE data, blocking until the client timeout (600s).
|
||||||
|
// with_line_timeout() only protects after the first line arrives, so
|
||||||
|
// users on older builds should set OLLAMA_STREAM_USAGE=false.
|
||||||
|
if !resolve_ollama_stream_usage() {
|
||||||
|
obj.remove("stream_options");
|
||||||
|
}
|
||||||
|
|
||||||
// Convert max_completion_tokens / max_tokens to Ollama's options.num_predict.
|
// Convert max_completion_tokens / max_tokens to Ollama's options.num_predict.
|
||||||
// Reasoning models emit max_completion_tokens; non-reasoning models emit max_tokens.
|
// Reasoning models emit max_completion_tokens; non-reasoning models emit max_tokens.
|
||||||
@@ -327,9 +353,34 @@ impl Provider for OllamaProvider {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
/// Per-chunk timeout for Ollama streaming responses.
|
/// Default per-chunk timeout for Ollama streaming responses (seconds).
|
||||||
/// If no new raw SSE data arrives within this duration, the connection is considered dead.
|
/// Configurable via OLLAMA_STREAM_TIMEOUT, GOOSE_STREAM_TIMEOUT, or falls back
|
||||||
const OLLAMA_CHUNK_TIMEOUT_SECS: u64 = 30;
|
/// to OLLAMA_TIMEOUT. Set high to accommodate slower models (CPU inference,
|
||||||
|
/// large parameter counts, complex reasoning).
|
||||||
|
const OLLAMA_DEFAULT_CHUNK_TIMEOUT_SECS: u64 = 120;
|
||||||
|
|
||||||
|
/// Resolve the per-chunk stream timeout from config.
|
||||||
|
/// Priority: OLLAMA_STREAM_TIMEOUT > GOOSE_STREAM_TIMEOUT > OLLAMA_TIMEOUT > default (120s).
|
||||||
|
/// Zero values are treated as invalid and skipped, since a zero timeout would
|
||||||
|
/// cause every chunk after the first to be treated as a stall.
|
||||||
|
fn resolve_ollama_chunk_timeout() -> u64 {
|
||||||
|
let config = crate::config::Config::global();
|
||||||
|
|
||||||
|
if let Ok(val) = config.get_param::<u64>("OLLAMA_STREAM_TIMEOUT") {
|
||||||
|
if val > 0 {
|
||||||
|
return val;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if let Ok(val) = config.get_param::<u64>("GOOSE_STREAM_TIMEOUT") {
|
||||||
|
if val > 0 {
|
||||||
|
return val;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
match config.get_param::<u64>("OLLAMA_TIMEOUT") {
|
||||||
|
Ok(val) if val > 0 => val,
|
||||||
|
_ => OLLAMA_DEFAULT_CHUNK_TIMEOUT_SECS,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
/// Wraps a line stream with a per-item timeout at the raw SSE level.
|
/// Wraps a line stream with a per-item timeout at the raw SSE level.
|
||||||
/// This detects dead connections without false-positive stalls during long
|
/// This detects dead connections without false-positive stalls during long
|
||||||
@@ -378,7 +429,8 @@ 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 timed_lines = with_line_timeout(framed, OLLAMA_CHUNK_TIMEOUT_SECS);
|
let chunk_timeout = resolve_ollama_chunk_timeout();
|
||||||
|
let timed_lines = with_line_timeout(framed, chunk_timeout);
|
||||||
let message_stream = response_to_streaming_message_ollama(timed_lines);
|
let message_stream = response_to_streaming_message_ollama(timed_lines);
|
||||||
pin!(message_stream);
|
pin!(message_stream);
|
||||||
|
|
||||||
@@ -450,7 +502,7 @@ mod tests {
|
|||||||
|
|
||||||
assert!(
|
assert!(
|
||||||
payload.get("stream_options").is_some(),
|
payload.get("stream_options").is_some(),
|
||||||
"create_request should produce stream_options (unsupported by Ollama)"
|
"create_request should produce stream_options for usage tracking"
|
||||||
);
|
);
|
||||||
assert!(
|
assert!(
|
||||||
payload.get("max_tokens").is_some(),
|
payload.get("max_tokens").is_some(),
|
||||||
@@ -459,11 +511,59 @@ mod tests {
|
|||||||
}
|
}
|
||||||
|
|
||||||
#[test]
|
#[test]
|
||||||
fn test_apply_ollama_options_strips_unsupported_fields() {
|
fn test_apply_ollama_options_preserves_stream_options_by_default() {
|
||||||
use crate::providers::formats::ollama::create_request;
|
use crate::providers::formats::ollama::create_request;
|
||||||
use crate::providers::utils::ImageFormat;
|
use crate::providers::utils::ImageFormat;
|
||||||
|
|
||||||
let _guard = env_lock::lock_env([("GOOSE_INPUT_LIMIT", None::<&str>)]);
|
let _guard = env_lock::lock_env([
|
||||||
|
("GOOSE_INPUT_LIMIT", None::<&str>),
|
||||||
|
("OLLAMA_STREAM_USAGE", 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_some(),
|
||||||
|
"stream_options should be preserved by default for usage tracking"
|
||||||
|
);
|
||||||
|
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");
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn test_apply_ollama_options_strips_stream_options_when_disabled() {
|
||||||
|
use crate::providers::formats::ollama::create_request;
|
||||||
|
use crate::providers::utils::ImageFormat;
|
||||||
|
|
||||||
|
let _guard = env_lock::lock_env([
|
||||||
|
("GOOSE_INPUT_LIMIT", None::<&str>),
|
||||||
|
("OLLAMA_STREAM_USAGE", Some("false")),
|
||||||
|
]);
|
||||||
let model_config = ModelConfig::new("llama3.1")
|
let model_config = ModelConfig::new("llama3.1")
|
||||||
.unwrap()
|
.unwrap()
|
||||||
.with_max_tokens(Some(4096));
|
.with_max_tokens(Some(4096));
|
||||||
@@ -483,74 +583,74 @@ mod tests {
|
|||||||
|
|
||||||
assert!(
|
assert!(
|
||||||
payload.get("stream_options").is_none(),
|
payload.get("stream_options").is_none(),
|
||||||
"stream_options should be removed for Ollama"
|
"stream_options should be removed when OLLAMA_STREAM_USAGE=false"
|
||||||
);
|
);
|
||||||
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]
|
#[test]
|
||||||
async fn test_stream_ollama_timeout_on_stall() {
|
fn test_resolve_ollama_chunk_timeout_defaults_to_ollama_timeout() {
|
||||||
use std::convert::Infallible;
|
let _guard = env_lock::lock_env([
|
||||||
|
("OLLAMA_STREAM_TIMEOUT", None::<&str>),
|
||||||
|
("GOOSE_STREAM_TIMEOUT", None::<&str>),
|
||||||
|
("OLLAMA_TIMEOUT", Some("300")),
|
||||||
|
]);
|
||||||
|
assert_eq!(resolve_ollama_chunk_timeout(), 300);
|
||||||
|
}
|
||||||
|
|
||||||
let (tx, rx) = tokio::sync::mpsc::channel::<Result<bytes::Bytes, Infallible>>(1);
|
#[test]
|
||||||
tx.send(Ok(bytes::Bytes::from(
|
fn test_resolve_ollama_chunk_timeout_prefers_stream_override() {
|
||||||
"data: {\"choices\":[{\"delta\":{\"content\":\"hi\"},\"index\":0}],\
|
let _guard = env_lock::lock_env([
|
||||||
\"model\":\"test\",\"object\":\"chat.completion.chunk\",\"created\":0}\n",
|
("OLLAMA_STREAM_TIMEOUT", Some("60")),
|
||||||
)))
|
("GOOSE_STREAM_TIMEOUT", Some("90")),
|
||||||
.await
|
("OLLAMA_TIMEOUT", Some("300")),
|
||||||
.unwrap();
|
]);
|
||||||
let stream = tokio_stream::wrappers::ReceiverStream::new(rx);
|
assert_eq!(resolve_ollama_chunk_timeout(), 60);
|
||||||
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(
|
#[test]
|
||||||
&ModelConfig::new("test").unwrap(),
|
fn test_resolve_ollama_chunk_timeout_uses_goose_stream_fallback() {
|
||||||
&json!({"model": "test"}),
|
let _guard = env_lock::lock_env([
|
||||||
)
|
("OLLAMA_STREAM_TIMEOUT", None::<&str>),
|
||||||
.unwrap();
|
("GOOSE_STREAM_TIMEOUT", Some("90")),
|
||||||
|
("OLLAMA_TIMEOUT", Some("300")),
|
||||||
|
]);
|
||||||
|
assert_eq!(resolve_ollama_chunk_timeout(), 90);
|
||||||
|
}
|
||||||
|
|
||||||
let mut msg_stream = stream_ollama(response, log).unwrap();
|
#[test]
|
||||||
|
fn test_resolve_ollama_chunk_timeout_uses_default_when_unset() {
|
||||||
|
let _guard = env_lock::lock_env([
|
||||||
|
("OLLAMA_STREAM_TIMEOUT", None::<&str>),
|
||||||
|
("GOOSE_STREAM_TIMEOUT", None::<&str>),
|
||||||
|
("OLLAMA_TIMEOUT", None::<&str>),
|
||||||
|
]);
|
||||||
|
assert_eq!(
|
||||||
|
resolve_ollama_chunk_timeout(),
|
||||||
|
OLLAMA_DEFAULT_CHUNK_TIMEOUT_SECS
|
||||||
|
);
|
||||||
|
}
|
||||||
|
|
||||||
let result =
|
#[test]
|
||||||
tokio::time::timeout(Duration::from_secs(OLLAMA_CHUNK_TIMEOUT_SECS + 5), async {
|
fn test_resolve_ollama_chunk_timeout_skips_zero_values() {
|
||||||
let mut last_err = None;
|
let _guard = env_lock::lock_env([
|
||||||
while let Some(item) = msg_stream.next().await {
|
("OLLAMA_STREAM_TIMEOUT", Some("0")),
|
||||||
if let Err(e) = item {
|
("GOOSE_STREAM_TIMEOUT", Some("0")),
|
||||||
last_err = Some(e);
|
("OLLAMA_TIMEOUT", Some("300")),
|
||||||
break;
|
]);
|
||||||
}
|
assert_eq!(resolve_ollama_chunk_timeout(), 300);
|
||||||
}
|
}
|
||||||
last_err
|
|
||||||
})
|
|
||||||
.await;
|
|
||||||
|
|
||||||
match result {
|
#[test]
|
||||||
Ok(Some(err)) => {
|
fn test_resolve_ollama_chunk_timeout_skips_all_zero_to_default() {
|
||||||
let err_msg = err.to_string();
|
let _guard = env_lock::lock_env([
|
||||||
assert!(
|
("OLLAMA_STREAM_TIMEOUT", Some("0")),
|
||||||
err_msg.contains("stream stalled"),
|
("GOOSE_STREAM_TIMEOUT", Some("0")),
|
||||||
"Expected stall timeout error, got: {}",
|
("OLLAMA_TIMEOUT", Some("0")),
|
||||||
err_msg
|
]);
|
||||||
);
|
assert_eq!(
|
||||||
}
|
resolve_ollama_chunk_timeout(),
|
||||||
Ok(None) => panic!("Expected timeout error but stream completed normally"),
|
OLLAMA_DEFAULT_CHUNK_TIMEOUT_SECS
|
||||||
Err(_) => panic!("Outer timeout elapsed -- per-chunk timeout did not fire"),
|
);
|
||||||
}
|
|
||||||
|
|
||||||
drop(tx);
|
|
||||||
}
|
}
|
||||||
|
|
||||||
#[test]
|
#[test]
|
||||||
|
|||||||
Reference in New Issue
Block a user