use std::sync::Arc;
use agent_base::llm_trait::ChatMessage;
use agent_base::{AgentResult, Middleware, PreLlmCtx};
use async_trait::async_trait;
use super::catalog::render_catalog;
use super::resolver::SkillResolver;
pub struct SkillCatalogRefreshMiddleware {
resolver: Arc<SkillResolver>,
}
impl SkillCatalogRefreshMiddleware {
pub fn new(resolver: Arc<SkillResolver>) -> Self {
Self { resolver }
}
}
#[async_trait]
impl Middleware for SkillCatalogRefreshMiddleware {
async fn on_pre_llm(&self, ctx: &mut PreLlmCtx) -> AgentResult<()> {
let Some(catalog) = render_catalog(&self.resolver) else {
return Ok(());
};
match ctx.messages.first_mut() {
Some(ChatMessage::System { content, .. }) => {
*content = super::catalog::refresh_catalog(content, &catalog);
}
_ => tracing::warn!(
"skill catalog refresh skipped: first message is not System \
(skills are loaded but the catalog cannot be refreshed)"
),
}
Ok(())
}
}
#[cfg(test)]
mod tests {
use super::*;
use agent_base::SessionId;
use std::path::Path;
const BASE: &str = "You are a coding agent.";
fn make_skill_dir(tmp: &Path, name: &str, body: &str) {
let dir = tmp.join("skills").join(name);
std::fs::create_dir_all(&dir).unwrap();
std::fs::write(
dir.join("SKILL.md"),
format!("---\nname: {name}\ndescription: test\nuser-invocable: true\n---\n\n{body}"),
)
.unwrap();
}
fn make_pre_llm_ctx(system_content: &str) -> PreLlmCtx {
PreLlmCtx {
session_id: SessionId {
id: 1,
external_id: None,
},
messages: vec![ChatMessage::System {
content: system_content.to_string(),
ephemeral: false,
}],
tools: vec![],
emit_fn: None,
turn_count: 1,
max_turns: 10,
}
}
fn stored_prompt(resolver: &SkillResolver) -> String {
format!("{BASE}\n\n{}", render_catalog(resolver).unwrap())
}
#[tokio::test]
async fn refresh_replaces_system_message_with_catalog() {
let tmp = tempfile::tempdir().unwrap();
make_skill_dir(tmp.path(), "alpha", "alpha body");
let resolver = Arc::new(SkillResolver::from_dirs(&[tmp.path().join("skills")]));
let mw = SkillCatalogRefreshMiddleware::new(resolver);
let mut ctx = make_pre_llm_ctx("original prompt");
mw.on_pre_llm(&mut ctx).await.unwrap();
match &ctx.messages[0] {
ChatMessage::System { content, .. } => {
assert!(
content.starts_with("original prompt"),
"base prompt must be preserved"
);
assert!(content.contains("## Skills"), "catalog must be appended");
assert!(content.contains("- alpha: test"), "skill must appear");
}
_ => panic!("expected System message"),
}
}
#[tokio::test]
async fn refresh_without_skills_keeps_original_prompt() {
let tmp = tempfile::tempdir().unwrap();
let resolver = Arc::new(SkillResolver::from_dirs(&[tmp.path().join("nonexistent")]));
let mw = SkillCatalogRefreshMiddleware::new(resolver);
let original = "base prompt only";
let mut ctx = make_pre_llm_ctx(original);
mw.on_pre_llm(&mut ctx).await.unwrap();
match &ctx.messages[0] {
ChatMessage::System { content, .. } => {
assert_eq!(content, original, "must be byte-identical without skills");
}
_ => panic!("expected System message"),
}
}
#[tokio::test]
async fn refresh_reflects_newly_added_skill() {
let tmp = tempfile::tempdir().unwrap();
let resolver = Arc::new(SkillResolver::from_dirs(&[tmp.path().join("skills")]));
let mw = SkillCatalogRefreshMiddleware::new(resolver.clone());
let mut ctx = make_pre_llm_ctx("prompt");
mw.on_pre_llm(&mut ctx).await.unwrap();
match &ctx.messages[0] {
ChatMessage::System { content, .. } => {
assert!(!content.contains("## Skills"));
}
_ => panic!(),
}
make_skill_dir(tmp.path(), "new-skill", "new body");
let resolver2 = Arc::new(SkillResolver::from_dirs(&[tmp.path().join("skills")]));
let mw2 = SkillCatalogRefreshMiddleware::new(resolver2);
let mut ctx = make_pre_llm_ctx("prompt");
mw2.on_pre_llm(&mut ctx).await.unwrap();
match &ctx.messages[0] {
ChatMessage::System { content, .. } => {
assert!(
content.contains("- new-skill: test"),
"new skill must appear"
);
}
_ => panic!(),
}
}
#[tokio::test]
async fn refresh_strips_builder_embedded_catalog_instead_of_duplicating() {
let tmp = tempfile::tempdir().unwrap();
make_skill_dir(tmp.path(), "alpha", "alpha body");
let resolver = Arc::new(SkillResolver::from_dirs(&[tmp.path().join("skills")]));
let stored = stored_prompt(&resolver);
assert!(stored.contains("## Skills"));
let mw = SkillCatalogRefreshMiddleware::new(resolver.clone());
let mut ctx = make_pre_llm_ctx(&stored);
mw.on_pre_llm(&mut ctx).await.unwrap();
match &ctx.messages[0] {
ChatMessage::System { content, .. } => {
assert_eq!(
content.matches("## Skills").count(),
1,
"catalog must appear exactly once, not duplicated per LLM call"
);
assert_eq!(
content,
&stored_prompt(&resolver),
"unchanged resolver → byte-identical recomposition"
);
}
_ => panic!("expected System message"),
}
}
#[tokio::test]
async fn refresh_preserves_trailing_sections_after_catalog() {
let tmp = tempfile::tempdir().unwrap();
make_skill_dir(tmp.path(), "alpha", "alpha body");
let resolver = Arc::new(SkillResolver::from_dirs(&[tmp.path().join("skills")]));
let suffix = "\n\n## Context Management\n\nuse history and notes tools.";
let stored = format!("{}{suffix}", stored_prompt(&resolver));
let mw = SkillCatalogRefreshMiddleware::new(resolver.clone());
let mut ctx = make_pre_llm_ctx(&stored);
mw.on_pre_llm(&mut ctx).await.unwrap();
match &ctx.messages[0] {
ChatMessage::System { content, .. } => {
assert_eq!(content.matches("## Skills").count(), 1, "{content}");
assert!(
content.ends_with(suffix),
"trailing section must survive: {content}"
);
assert_eq!(
content, &stored,
"unchanged resolver + suffix → byte-identical to builder output"
);
}
_ => panic!("expected System message"),
}
}
#[tokio::test]
async fn refresh_is_idempotent_across_repeated_calls() {
let tmp = tempfile::tempdir().unwrap();
make_skill_dir(tmp.path(), "alpha", "alpha body");
let resolver = Arc::new(SkillResolver::from_dirs(&[tmp.path().join("skills")]));
let mw = SkillCatalogRefreshMiddleware::new(resolver.clone());
let mut ctx = make_pre_llm_ctx(&stored_prompt(&resolver));
mw.on_pre_llm(&mut ctx).await.unwrap();
let after_first = match &ctx.messages[0] {
ChatMessage::System { content, .. } => content.clone(),
_ => panic!(),
};
let mut ctx2 = make_pre_llm_ctx(&after_first);
mw.on_pre_llm(&mut ctx2).await.unwrap();
match &ctx2.messages[0] {
ChatMessage::System { content, .. } => {
assert_eq!(
content, &after_first,
"repeated refresh must be a fixed point"
);
assert_eq!(content.matches("## Skills").count(), 1);
}
_ => panic!("expected System message"),
}
}
#[tokio::test]
async fn non_system_first_message_left_untouched() {
let tmp = tempfile::tempdir().unwrap();
make_skill_dir(tmp.path(), "alpha", "alpha body");
let resolver = Arc::new(SkillResolver::from_dirs(&[tmp.path().join("skills")]));
let mw = SkillCatalogRefreshMiddleware::new(resolver);
let mut ctx = make_pre_llm_ctx("unused");
ctx.messages = vec![ChatMessage::User {
content: "hi".to_string(),
images: vec![],
ephemeral: false,
}];
mw.on_pre_llm(&mut ctx).await.unwrap();
match &ctx.messages[0] {
ChatMessage::User { content, .. } => {
assert_eq!(content, "hi", "user message must survive")
}
other => panic!("expected User message, got {other:?}"),
}
}
#[tokio::test]
async fn empty_messages_do_not_panic() {
let tmp = tempfile::tempdir().unwrap();
make_skill_dir(tmp.path(), "alpha", "alpha body");
let resolver = Arc::new(SkillResolver::from_dirs(&[tmp.path().join("skills")]));
let mw = SkillCatalogRefreshMiddleware::new(resolver);
let mut ctx = make_pre_llm_ctx("unused");
ctx.messages.clear();
mw.on_pre_llm(&mut ctx).await.unwrap();
assert!(ctx.messages.is_empty());
}
}