Surface resolved Databricks model metadata (#9206)

Signed-off-by: jh-block <jhugo@block.xyz>
This commit is contained in:
jh-block
2026-05-20 12:50:44 +02:00
committed by GitHub
parent 98a54e9ec6
commit 8f9270e295
19 changed files with 369 additions and 47 deletions
+5 -3
View File
@@ -20,9 +20,10 @@ use goose::config::declarative_providers::{
DeclarativeProviderConfig, EnvVarConfig, LoadedProvider, ProviderEngine,
};
use goose::conversation::message::{
ActionRequired, ActionRequiredData, FrontendToolRequest, Message, MessageContent,
MessageMetadata, RedactedThinkingContent, SystemNotificationContent, SystemNotificationType,
ThinkingContent, TokenState, ToolConfirmationRequest, ToolRequest, ToolResponse,
ActionRequired, ActionRequiredData, FrontendToolRequest, InferenceMetadata, Message,
MessageContent, MessageMetadata, RedactedThinkingContent, SystemNotificationContent,
SystemNotificationType, ThinkingContent, TokenState, ToolConfirmationRequest, ToolRequest,
ToolResponse,
};
use crate::routes::recipe_utils::RecipeManifest;
@@ -528,6 +529,7 @@ derive_utoipa!(IconTheme as IconThemeSchema);
Message,
MessageContent,
MessageMetadata,
InferenceMetadata,
TokenState,
ContentSchema,
EmbeddedResourceSchema,
+36 -2
View File
@@ -34,8 +34,8 @@ use crate::context_mgmt::{
check_if_compaction_needed, compact_messages, DEFAULT_COMPACTION_THRESHOLD,
};
use crate::conversation::message::{
ActionRequiredData, Message, MessageContent, ProviderMetadata, SystemNotificationType,
ToolRequest,
ActionRequiredData, InferenceMetadata, Message, MessageContent, ProviderMetadata,
SystemNotificationType, ToolRequest,
};
use crate::conversation::{debug_conversation_fix, fix_conversation, Conversation};
use crate::mcp_utils::ToolResult;
@@ -1586,9 +1586,22 @@ impl Agent {
self.reset_retry_attempts().await;
let provider = self.provider().await?;
let provider_name = provider.get_name().to_string();
let requested_model = provider.get_model_config().model_name;
let inference = provider
.fetch_model_info(&requested_model)
.await
.ok()
.and_then(|model_info| model_info.resolved_model)
.map(|resolved_model| InferenceMetadata {
provider: provider_name,
requested_model,
resolved_model: Some(resolved_model),
});
let session_manager = self.config.session_manager.clone();
let session_id = session_config.id.clone();
if !self.config.disable_session_naming {
let provider = provider.clone();
let manager_for_spawn = session_manager.clone();
let session_name_update_tx = self.config.session_name_update_tx.clone();
tokio::spawn(async move {
@@ -1726,6 +1739,17 @@ impl Agent {
)
.await;
let filtered_response = if let Some(inference) = inference.as_ref() {
filtered_response.with_inference(inference.clone())
} else {
filtered_response
};
let response = if let Some(inference) = inference.as_ref() {
response.with_inference(inference.clone())
} else {
response
};
surfaced_thinking_in_turn |= filtered_response.content.iter().any(
|content| {
matches!(
@@ -2234,6 +2258,16 @@ impl Agent {
}
}
let messages_to_add = if let Some(ref inference) = inference {
Conversation::new_unvalidated(
messages_to_add
.into_iter()
.map(|message| message.with_inference_if_assistant(inference.clone())),
)
} else {
messages_to_add
};
for msg in &messages_to_add {
session_manager.add_message(&session_config.id, msg).await?;
}
+1 -1
View File
@@ -143,7 +143,7 @@ pub async fn compact_messages(
// This is the most recent message and we're preserving it by adding a fresh copy
MessageMetadata::invisible()
} else {
msg.metadata.with_agent_invisible()
msg.metadata.clone().with_agent_invisible()
};
let updated_msg = msg.clone().with_metadata(updated_metadata);
final_messages.push(updated_msg);
+35 -2
View File
@@ -641,14 +641,25 @@ impl From<PromptMessage> for Message {
}
}
#[derive(ToSchema, Clone, Copy, PartialEq, Serialize, Deserialize, Debug)]
/// Metadata for message visibility
#[derive(ToSchema, Clone, PartialEq, Serialize, Deserialize, Debug)]
#[serde(rename_all = "camelCase")]
pub struct InferenceMetadata {
pub provider: String,
pub requested_model: String,
#[serde(skip_serializing_if = "Option::is_none")]
pub resolved_model: Option<String>,
}
#[derive(ToSchema, Clone, PartialEq, Serialize, Deserialize, Debug)]
/// Metadata for message visibility and model inference details
#[serde(rename_all = "camelCase")]
pub struct MessageMetadata {
/// Whether the message should be visible to the user in the UI
pub user_visible: bool,
/// Whether the message should be included in the agent's context window
pub agent_visible: bool,
#[serde(skip_serializing_if = "Option::is_none")]
pub inference: Option<InferenceMetadata>,
}
impl Default for MessageMetadata {
@@ -656,6 +667,7 @@ impl Default for MessageMetadata {
MessageMetadata {
user_visible: true,
agent_visible: true,
inference: None,
}
}
}
@@ -666,6 +678,7 @@ impl MessageMetadata {
MessageMetadata {
user_visible: false,
agent_visible: true,
..Default::default()
}
}
@@ -674,6 +687,7 @@ impl MessageMetadata {
MessageMetadata {
user_visible: true,
agent_visible: false,
..Default::default()
}
}
@@ -682,6 +696,7 @@ impl MessageMetadata {
MessageMetadata {
user_visible: false,
agent_visible: false,
..Default::default()
}
}
@@ -716,6 +731,11 @@ impl MessageMetadata {
..self
}
}
pub fn with_inference(mut self, inference: InferenceMetadata) -> Self {
self.inference = Some(inference);
self
}
}
#[derive(ToSchema, Clone, PartialEq, Serialize, Deserialize, Debug)]
@@ -996,6 +1016,19 @@ impl Message {
self
}
pub fn with_inference(mut self, inference: InferenceMetadata) -> Self {
self.metadata = self.metadata.with_inference(inference);
self
}
pub fn with_inference_if_assistant(self, inference: InferenceMetadata) -> Self {
if self.role == Role::Assistant && self.metadata.inference.is_none() {
self.with_inference(inference)
} else {
self
}
}
pub fn user_only(mut self) -> Self {
self.metadata.user_visible = true;
self.metadata.agent_visible = false;
+15
View File
@@ -43,11 +43,26 @@ impl Conversation {
}
pub fn push(&mut self, message: Message) {
if message.content.is_empty() && message.metadata.inference.is_some() {
if let Some(existing) = self
.0
.iter_mut()
.rev()
.find(|m| m.role == message.role && m.is_user_visible())
{
existing.metadata.inference = message.metadata.inference;
}
return;
}
if let Some(last) = self
.0
.last_mut()
.filter(|m| m.id.is_some() && m.id == message.id)
{
if message.metadata.inference.is_some() {
last.metadata.inference = message.metadata.inference.clone();
}
match (last.content.last_mut(), message.content.last()) {
(Some(MessageContent::Text(ref mut last)), Some(MessageContent::Text(new)))
if message.content.len() == 1 =>
+9
View File
@@ -385,6 +385,9 @@ pub static MSG_COUNT_FOR_SESSION_NAME_GENERATION: usize = 3;
pub struct ModelInfo {
/// The name of the model
pub name: String,
/// The underlying model resolved from provider metadata, when the configured model is an alias or endpoint.
#[serde(default, skip_serializing_if = "Option::is_none")]
pub resolved_model: Option<String>,
/// The maximum context length this model supports
pub context_limit: usize,
/// Cost per token for input in USD (optional)
@@ -405,6 +408,7 @@ impl ModelInfo {
pub fn new(name: impl Into<String>, context_limit: usize) -> Self {
Self {
name: name.into(),
resolved_model: None,
context_limit,
input_token_cost: None,
output_token_cost: None,
@@ -423,6 +427,7 @@ impl ModelInfo {
) -> Self {
Self {
name: name.into(),
resolved_model: None,
context_limit,
input_token_cost: Some(input_cost),
output_token_cost: Some(output_cost),
@@ -448,6 +453,7 @@ fn model_info_for_provider_model(provider_name: &str, model_name: &str) -> Model
ModelInfo {
name: model_name.to_string(),
resolved_model: None,
context_limit: ModelConfig::new_or_fail(model_name)
.with_canonical_limits(provider_name)
.context_limit(),
@@ -1778,6 +1784,7 @@ mod tests {
// Test direct ModelInfo creation
let info = ModelInfo {
name: "test-model".to_string(),
resolved_model: None,
context_limit: 1000,
input_token_cost: None,
output_token_cost: None,
@@ -1790,6 +1797,7 @@ mod tests {
// Test equality
let info2 = ModelInfo {
name: "test-model".to_string(),
resolved_model: None,
context_limit: 1000,
input_token_cost: None,
output_token_cost: None,
@@ -1802,6 +1810,7 @@ mod tests {
// Test inequality
let info3 = ModelInfo {
name: "test-model".to_string(),
resolved_model: None,
context_limit: 2000,
input_token_cost: None,
output_token_cost: None,
+9
View File
@@ -539,6 +539,7 @@ impl DatabricksProvider {
ModelInfo {
name: info.name,
resolved_model: info.upstream_model_name,
context_limit,
input_token_cost: None,
output_token_cost: None,
@@ -999,6 +1000,14 @@ mod tests {
assert_eq!(info.name, "goose");
assert_eq!(info.upstream_model_name.as_deref(), Some("claude-opus-4.6"));
assert_eq!(info.reasoning, Some(true));
let model_info = DatabricksProvider::model_info_from_endpoint(info);
assert_eq!(model_info.name, "goose");
assert_eq!(
model_info.resolved_model.as_deref(),
Some("claude-opus-4.6")
);
assert!(model_info.reasoning);
}
#[test]
+25 -7
View File
@@ -734,15 +734,22 @@ pub fn get_usage(usage: &Value) -> Usage {
.with_cache_tokens(cache_read_input_tokens, cache_write_input_tokens)
}
fn extract_usage_with_output_tokens(chunk: &StreamingChunk) -> Option<ProviderUsage> {
fn extract_usage_with_output_tokens(
chunk: &StreamingChunk,
fallback_model: Option<&str>,
) -> Option<ProviderUsage> {
chunk
.usage
.as_ref()
.and_then(|u| {
chunk.model.as_ref().map(|model| ProviderUsage {
usage: get_usage(u),
model: model.clone(),
})
chunk
.model
.as_deref()
.or(fallback_model)
.map(|model| ProviderUsage {
usage: get_usage(u),
model: model.to_string(),
})
})
.filter(|u| u.usage.output_tokens.is_some())
}
@@ -901,6 +908,7 @@ where
// reasoning will arrive. Emitting it immediately and then receiving
// reasoning_content in a later chunk would produce duplicated reasoning.
let mut pending_inline_thinking = String::new();
let mut last_seen_model: Option<String> = None;
'outer: while let Some(response) = stream.next().await {
let response_str = response?;
@@ -917,6 +925,9 @@ where
let chunk: StreamingChunk = parse_streaming_chunk(
line.ok_or_else(|| anyhow!("unexpected stream format"))?
)?;
if let Some(model) = &chunk.model {
last_seen_model = Some(model.clone());
}
if !chunk.choices.is_empty() {
if let Some(details) = &chunk.choices[0].delta.reasoning_details {
@@ -931,7 +942,7 @@ where
}
}
let mut usage = extract_usage_with_output_tokens(&chunk);
let mut usage = extract_usage_with_output_tokens(&chunk, last_seen_model.as_deref());
if chunk.choices.is_empty() {
yield (None, usage)
@@ -959,8 +970,11 @@ where
}
let tool_chunk: StreamingChunk = parse_streaming_chunk(line)?;
if let Some(model) = &tool_chunk.model {
last_seen_model = Some(model.clone());
}
if let Some(chunk_usage) = extract_usage_with_output_tokens(&tool_chunk) {
if let Some(chunk_usage) = extract_usage_with_output_tokens(&tool_chunk, last_seen_model.as_deref()) {
usage = Some(chunk_usage);
}
@@ -2528,6 +2542,10 @@ data: [DONE]
assert_eq!(result.tool_calls.len(), 1, "Expected 1 tool call");
assert_eq!(result.tool_calls[0], "developer__shell");
assert_usage_yielded_once(&result, 8320, 172, 8492);
assert_eq!(
result.usage.as_ref().map(|usage| usage.model.as_str()),
Some("gpt-5.2-1106-preview")
);
Ok(())
}
@@ -177,6 +177,7 @@ impl ProviderRegistry {
.iter()
.map(|m| ModelInfo {
name: m.name.clone(),
resolved_model: None,
context_limit: m.context_limit,
input_token_cost: m.input_token_cost,
output_token_cost: m.output_token_cost,