mirror of
https://github.com/fabro-sh/fabro.git
synced 2026-10-07 03:00:29 +00:00
fix(llm): require explicit model test client
This commit is contained in:
parent
ec3d65928f
commit
7858a73146
2 changed files with 12 additions and 38 deletions
|
|
@ -1,13 +1,11 @@
|
|||
use std::sync::Arc;
|
||||
use std::time::Duration;
|
||||
|
||||
use fabro_auth::EnvCredentialSource;
|
||||
use fabro_model::Model;
|
||||
use strum::{EnumString, IntoStaticStr};
|
||||
use tokio::time;
|
||||
|
||||
use crate::client::Client;
|
||||
use crate::error::Error;
|
||||
use crate::generate::{self, GenerateParams};
|
||||
use crate::tools::Tool;
|
||||
use crate::types::{GenerateResult, ReasoningEffort};
|
||||
|
|
@ -73,22 +71,10 @@ impl ModelTestOutcome {
|
|||
}
|
||||
}
|
||||
|
||||
pub async fn run_model_test(info: &Model, mode: ModelTestMode) -> ModelTestOutcome {
|
||||
run_model_test_inner(info, mode, None).await
|
||||
}
|
||||
|
||||
pub async fn run_model_test_with_client(
|
||||
pub async fn run_model_test(
|
||||
info: &Model,
|
||||
mode: ModelTestMode,
|
||||
client: Arc<Client>,
|
||||
) -> ModelTestOutcome {
|
||||
run_model_test_inner(info, mode, Some(client)).await
|
||||
}
|
||||
|
||||
async fn run_model_test_inner(
|
||||
info: &Model,
|
||||
mode: ModelTestMode,
|
||||
client: Option<Arc<Client>>,
|
||||
) -> ModelTestOutcome {
|
||||
match mode {
|
||||
ModelTestMode::Basic => run_basic_test(info, client).await,
|
||||
|
|
@ -96,12 +82,7 @@ async fn run_model_test_inner(
|
|||
}
|
||||
}
|
||||
|
||||
async fn run_basic_test(info: &Model, client: Option<Arc<Client>>) -> ModelTestOutcome {
|
||||
let client = match resolve_client(client).await {
|
||||
Ok(client) => client,
|
||||
Err(err) => return ModelTestOutcome::error(err.to_string()),
|
||||
};
|
||||
|
||||
async fn run_basic_test(info: &Model, client: Arc<Client>) -> ModelTestOutcome {
|
||||
let params = GenerateParams::new(&info.id, client)
|
||||
.provider(info.provider.as_str())
|
||||
.prompt("Say OK")
|
||||
|
|
@ -120,11 +101,7 @@ async fn run_basic_test(info: &Model, client: Option<Arc<Client>>) -> ModelTestO
|
|||
}
|
||||
}
|
||||
|
||||
async fn run_deep_test(info: &Model, client: Option<Arc<Client>>) -> ModelTestOutcome {
|
||||
let client = match resolve_client(client).await {
|
||||
Ok(client) => client,
|
||||
Err(err) => return ModelTestOutcome::error(err.to_string()),
|
||||
};
|
||||
async fn run_deep_test(info: &Model, client: Arc<Client>) -> ModelTestOutcome {
|
||||
let Some(params) = build_deep_test_params(info, client) else {
|
||||
return ModelTestOutcome::error("model does not support tools");
|
||||
};
|
||||
|
|
@ -191,15 +168,6 @@ fn build_deep_test_params(info: &Model, client: Arc<Client>) -> Option<GenerateP
|
|||
Some(params)
|
||||
}
|
||||
|
||||
async fn resolve_client(client: Option<Arc<Client>>) -> Result<Arc<Client>, Error> {
|
||||
if let Some(client) = client {
|
||||
return Ok(client);
|
||||
}
|
||||
|
||||
let source = EnvCredentialSource::new();
|
||||
Client::from_source(&source).await
|
||||
}
|
||||
|
||||
fn validate_deep_result(result: &GenerateResult) -> Result<(), String> {
|
||||
if result.steps.len() < 2 {
|
||||
return Err("model did not call tool".to_string());
|
||||
|
|
@ -218,6 +186,8 @@ fn validate_deep_result(result: &GenerateResult) -> Result<(), String> {
|
|||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use std::collections::HashMap;
|
||||
|
||||
use fabro_model::{ModelCosts, ModelFeatures, ModelLimits, Provider};
|
||||
|
||||
use super::*;
|
||||
|
|
@ -261,6 +231,10 @@ mod tests {
|
|||
}
|
||||
}
|
||||
|
||||
fn empty_test_client() -> Arc<Client> {
|
||||
Arc::new(Client::new(HashMap::new(), None, vec![]))
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn run_model_test_deep_errors_when_model_lacks_tools() {
|
||||
let info = test_model_with(ModelFeatures {
|
||||
|
|
@ -270,7 +244,7 @@ mod tests {
|
|||
effort: true,
|
||||
});
|
||||
|
||||
let outcome = run_model_test(&info, ModelTestMode::Deep).await;
|
||||
let outcome = run_model_test(&info, ModelTestMode::Deep, empty_test_client()).await;
|
||||
|
||||
assert_eq!(outcome.status, ModelTestStatus::Error);
|
||||
assert_eq!(
|
||||
|
|
|
|||
|
|
@ -47,7 +47,7 @@ use fabro_interview::{
|
|||
};
|
||||
use fabro_llm::client::Client as LlmClient;
|
||||
use fabro_llm::generate::{GenerateParams, generate_object};
|
||||
use fabro_llm::model_test::{ModelTestMode, run_model_test_with_client};
|
||||
use fabro_llm::model_test::{ModelTestMode, run_model_test};
|
||||
use fabro_llm::types::{
|
||||
ContentPart, FinishReason, Message as LlmMessage, Request as LlmRequest, Role, ToolChoice,
|
||||
ToolDefinition,
|
||||
|
|
@ -6685,7 +6685,7 @@ async fn test_model(
|
|||
}
|
||||
let client = Arc::new(llm_result.client);
|
||||
|
||||
let outcome = run_model_test_with_client(info, mode, client).await;
|
||||
let outcome = run_model_test(info, mode, client).await;
|
||||
Json(serde_json::json!({
|
||||
"model_id": info.id,
|
||||
"status": outcome.status.as_str(),
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue