Surface resolved Databricks model metadata (#9206)
Signed-off-by: jh-block <jhugo@block.xyz>
This commit is contained in:
@@ -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,
|
||||
|
||||
@@ -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?;
|
||||
}
|
||||
|
||||
@@ -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);
|
||||
|
||||
@@ -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;
|
||||
|
||||
@@ -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 =>
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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]
|
||||
|
||||
@@ -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,
|
||||
|
||||
Reference in New Issue
Block a user