Skip to content

Commit 52f8949

Browse files
senamakelmedullabot
andcommitted
fix(harness): reserve hosted terminal progress
Co-authored-by: Medulla <medulla@tinyhumans.ai>
1 parent 2bd5cde commit 52f8949

4 files changed

Lines changed: 81 additions & 19 deletions

File tree

‎crates/tinyagents-harness/src/agent_loop/model_call.rs‎

Lines changed: 38 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -826,20 +826,51 @@ impl<State: Send + Sync, Ctx: Send + Sync> ModelCallBase<'_, State, Ctx> {
826826
/// fail-closed behaviour `run_loop` already has for a pre-wrap override.
827827
/// Silently substituting a different model is the one outcome that is
828828
/// never acceptable.
829-
fn rebind(
829+
async fn rebind(
830830
&self,
831831
ctx: &mut RunContext<Ctx>,
832832
request: &ModelRequest,
833-
) -> ResolvedModelBinding<State> {
833+
) -> Result<ResolvedModelBinding<State>> {
834834
let captured = || ResolvedModelBinding {
835835
resolved: self.resolved.clone(),
836836
model: Arc::clone(&self.model),
837837
};
838838
let Some(requested) = request.model.as_deref() else {
839-
return captured();
839+
return Ok(captured());
840840
};
841+
if let Some(host_run) = self.harness.host_run_binding(ctx.instance_id())? {
842+
let mut resolve = crate::host::ModelResolveRequest::new(host_run.agent_id.clone());
843+
if ctx.depth() == 0 {
844+
resolve = resolve.as_team_lead();
845+
}
846+
if let Some(role) = host_run.role.clone() {
847+
resolve = resolve.with_role(role);
848+
}
849+
if let Some(pin) = request.model.clone().or(host_run.model_pin.clone()) {
850+
resolve = resolve.with_model_pin(pin);
851+
}
852+
if let Some(capabilities) = request.required_capabilities.clone() {
853+
resolve = resolve.with_required_capabilities(capabilities);
854+
}
855+
let model = host_run.host.models.resolve(&resolve).await.map_err(|error| {
856+
tinyagents_tracing::warn!(%error, agent_id = %host_run.agent_id, "[host] wrap model resolution failed");
857+
TinyAgentsError::Model("host model resolution failed".to_string())
858+
})?;
859+
let name = model
860+
.profile()
861+
.and_then(|profile| profile.model.clone())
862+
.unwrap_or_else(|| format!("host:{}", host_run.agent_id));
863+
return Ok(ResolvedModelBinding {
864+
resolved: ResolvedModel {
865+
name,
866+
requested: resolve.model_pin,
867+
source: ModelResolutionSource::RequestOverride,
868+
},
869+
model,
870+
});
871+
}
841872
if requested == self.resolved.name {
842-
return captured();
873+
return Ok(captured());
843874
}
844875
match self.harness.models.resolve_request(request, None, None) {
845876
Some(binding)
@@ -852,7 +883,7 @@ impl<State: Send + Sync, Ctx: Send + Sync> ModelCallBase<'_, State, Ctx> {
852883
to = %binding.resolved.name,
853884
"[model] wrap layer overrode the model; re-resolved the binding"
854885
);
855-
binding
886+
Ok(binding)
856887
}
857888
_ => {
858889
tinyagents_tracing::warn!(
@@ -865,7 +896,7 @@ impl<State: Send + Sync, Ctx: Send + Sync> ModelCallBase<'_, State, Ctx> {
865896
requested: requested.to_string(),
866897
resolved: self.resolved.name.clone(),
867898
});
868-
captured()
899+
Ok(captured())
869900
}
870901
}
871902
}
@@ -881,7 +912,7 @@ impl<State: Send + Sync, Ctx: Send + Sync> ModelBaseCall<State, Ctx>
881912
request: ModelRequest,
882913
) -> BoxModelFuture<'a> {
883914
Box::pin(async move {
884-
let binding = self.rebind(ctx, &request);
915+
let binding = self.rebind(ctx, &request).await?;
885916
self.harness
886917
.invoke_model_with_retry(
887918
state,

‎crates/tinyagents-harness/src/runtime/agent.rs‎

Lines changed: 36 additions & 11 deletions
Original file line numberDiff line numberDiff line change
@@ -128,7 +128,28 @@ struct PreparedAgentTurn<State: Send + Sync> {
128128
run_id: crate::ids::RunId,
129129
input_text: String,
130130
messages: Vec<tinyinference_llm::message::Message>,
131-
progress: Option<tokio::sync::mpsc::Sender<ProgressEvent>>,
131+
progress: Option<ProgressSender>,
132+
}
133+
134+
#[derive(Clone)]
135+
pub(crate) struct ProgressSender {
136+
tx: tokio::sync::mpsc::Sender<ProgressEvent>,
137+
nonterminal_slots: std::sync::Arc<tokio::sync::Semaphore>,
138+
}
139+
140+
impl ProgressSender {
141+
fn send_nonterminal(&self, event: ProgressEvent) {
142+
let Ok(permit) = self.nonterminal_slots.clone().try_acquire_owned() else {
143+
return;
144+
};
145+
if self.tx.try_send(event).is_ok() {
146+
permit.forget();
147+
}
148+
}
149+
150+
fn send_terminal(&self, event: ProgressEvent) {
151+
let _ = self.tx.try_send(event);
152+
}
132153
}
133154

134155
impl<State: Send + Sync> Clone for PreparedAgentTurn<State> {
@@ -450,11 +471,7 @@ impl<State: Send + Sync, Ctx: Send + Sync> AgentHarness<State, Ctx> {
450471
let Some(progress) = binding.progress else {
451472
return;
452473
};
453-
if progress.try_send(event).is_err() {
454-
tinyagents_tracing::debug!(
455-
"[host] dropping progress event because the bounded queue is full or closed"
456-
);
457-
}
474+
progress.send_nonterminal(event);
458475
}
459476
}
460477

@@ -484,12 +501,12 @@ async fn finish_host_turn<State: Send + Sync>(
484501
// consumer completely outside the agent's critical path.
485502
if let Some(progress) = &prepared.progress {
486503
if let Some(message) = error {
487-
let _ = progress.try_send(ProgressEvent::Error {
504+
progress.send_terminal(ProgressEvent::Error {
488505
run: prepared.run_id.clone(),
489506
message,
490507
});
491508
} else {
492-
let _ = progress.try_send(ProgressEvent::Finished {
509+
progress.send_terminal(ProgressEvent::Finished {
493510
run: prepared.run_id.clone(),
494511
usage: Some(run.usage.usage),
495512
});
@@ -533,21 +550,29 @@ async fn finish_host_turn<State: Send + Sync>(
533550

534551
fn start_progress_dispatcher(
535552
sink: Option<std::sync::Arc<dyn crate::host::ProgressSink>>,
536-
) -> Option<tokio::sync::mpsc::Sender<ProgressEvent>> {
553+
) -> Option<ProgressSender> {
537554
let sink = sink?;
538555
let handle = tokio::runtime::Handle::try_current().ok()?;
539556
// Progress is observational. Bound it so a slow sink cannot retain every
540557
// streamed token; producers use `try_send` and drop overflowed updates.
541558
// One slot is reserved for the single terminal outcome. Producers use
542559
// `try_send` for ordinary progress, so at most 128 nonterminal events can
543560
// fill before finalization claims the remaining slot.
544-
let (tx, mut rx) = tokio::sync::mpsc::channel(129);
561+
let (tx, mut rx) = tokio::sync::mpsc::channel::<ProgressEvent>(129);
562+
let nonterminal_slots = std::sync::Arc::new(tokio::sync::Semaphore::new(128));
563+
let released_slots = nonterminal_slots.clone();
545564
handle.spawn(async move {
546565
while let Some(event) = rx.recv().await {
566+
if !event.is_terminal() {
567+
released_slots.add_permits(1);
568+
}
547569
sink.emit(event).await;
548570
}
549571
});
550-
Some(tx)
572+
Some(ProgressSender {
573+
tx,
574+
nonterminal_slots,
575+
})
551576
}
552577

553578
async fn screen_user_messages<State: Send + Sync>(

‎crates/tinyagents-harness/src/runtime/test.rs‎

Lines changed: 6 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -49,6 +49,7 @@ struct RetryableClassifier;
4949
struct LeadRecordingResolver {
5050
model: Arc<dyn ChatModel<()>>,
5151
team_lead_flags: Mutex<Vec<bool>>,
52+
model_pins: Mutex<Vec<Option<String>>>,
5253
}
5354

5455
struct RecordingBudget {
@@ -146,6 +147,10 @@ impl crate::host::ModelResolver<()> for LeadRecordingResolver {
146147
.lock()
147148
.expect("resolver lock")
148149
.push(request.is_team_lead);
150+
self.model_pins
151+
.lock()
152+
.expect("resolver lock")
153+
.push(request.model_pin.clone());
149154
Ok(Arc::clone(&self.model))
150155
}
151156
}
@@ -524,6 +529,7 @@ async fn hosted_model_resolution_marks_only_root_contexts_as_team_leads() {
524529
let resolver = Arc::new(LeadRecordingResolver {
525530
model: model.clone(),
526531
team_lead_flags: Mutex::new(Vec::new()),
532+
model_pins: Mutex::new(Vec::new()),
527533
});
528534
let host = crate::host::HostCapabilities::new(
529535
Arc::new(StaticContextComposer::empty()),

‎crates/tinyagents-harness/src/runtime/types.rs‎

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -50,7 +50,7 @@ pub(crate) struct HostRunBinding<State: Send + Sync> {
5050
/// list is a host boundary enforced for schemas and dispatch alike.
5151
pub(crate) allowed_tools: HashSet<String>,
5252
/// Per-turn ordered, nonblocking projection to the optional progress sink.
53-
pub(crate) progress: Option<tokio::sync::mpsc::Sender<crate::host::ProgressEvent>>,
53+
pub(crate) progress: Option<super::agent::ProgressSender>,
5454
}
5555

5656
impl<State: Send + Sync> Clone for HostRunBinding<State> {

0 commit comments

Comments
 (0)