use std::sync::{Arc, Mutex};
use std::time::Instant;
use rmcp::model::{CompleteRequestParams, CompleteResult, CompletionInfo, Reference};
use crate::context::ToolContext;
use crate::gating::Gate;
use crate::meta::Surface;
pub const DEVICE_TEMPLATE: &str = "tailnet://device/{device_id}";
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
enum Slot {
Device,
Peer,
Subject,
}
fn slot_for(reference: &Reference, argument: &str) -> Option<Slot> {
match reference {
Reference::Resource(template) if template.uri == DEVICE_TEMPLATE => {
(argument == "device_id").then_some(Slot::Device)
}
Reference::Prompt(prompt) => match (prompt.name.as_str(), argument) {
("diagnose_connectivity", "peer") => Some(Slot::Peer),
("audit_tailnet_access", "subject") => Some(Slot::Subject),
_ => None,
},
Reference::Resource(_) => None,
_ => None,
}
}
impl Slot {
const fn surface(self) -> Surface {
match self {
Self::Device | Self::Subject => Surface::Tailnet,
Self::Peer => Surface::Local,
}
}
}
#[derive(Clone, Copy, Debug, PartialEq, Eq, PartialOrd, Ord)]
enum Kind {
User,
Tag,
Device,
}
#[derive(Debug)]
struct Candidate {
value: String,
kind: Kind,
known_as: Vec<String>,
}
impl Candidate {
fn new(value: impl Into<String>, kind: Kind, also: impl IntoIterator<Item = String>) -> Self {
let value = value.into();
let mut known_as = vec![normalise(&value)];
known_as.extend(
also.into_iter()
.map(|other| normalise(&other))
.filter(|other| !other.is_empty()),
);
known_as.sort_unstable();
known_as.dedup();
Self {
value,
kind,
known_as,
}
}
fn rank(&self, typed: &str) -> Option<u8> {
if typed.is_empty() {
return Some(2);
}
self.known_as
.iter()
.filter_map(|known| {
if known == typed {
Some(0)
} else if known.starts_with(typed) {
Some(1)
} else if known.contains(typed) {
Some(2)
} else {
None
}
})
.min()
}
}
fn normalise(name: &str) -> String {
name.trim().trim_end_matches('.').to_ascii_lowercase()
}
#[derive(Debug)]
struct Bucket {
tokens: f64,
last: Instant,
}
#[derive(Clone, Debug)]
pub struct Limiter {
bucket: Arc<Mutex<Bucket>>,
}
impl Default for Limiter {
fn default() -> Self {
Self::new()
}
}
impl Limiter {
const BURST: f64 = 20.0;
const PER_SECOND: f64 = 20.0;
pub fn new() -> Self {
Self {
bucket: Arc::new(Mutex::new(Bucket {
tokens: Self::BURST,
last: Instant::now(),
})),
}
}
fn allow(&self) -> bool {
let Ok(mut bucket) = self.bucket.lock() else {
return true;
};
let now = Instant::now();
let earned = now.duration_since(bucket.last).as_secs_f64() * Self::PER_SECOND;
bucket.tokens = (bucket.tokens + earned).min(Self::BURST);
bucket.last = now;
if bucket.tokens < 1.0 {
return false;
}
bucket.tokens -= 1.0;
true
}
}
pub async fn complete(
ctx: &ToolContext,
gate: &Gate,
limiter: &Limiter,
request: &CompleteRequestParams,
) -> CompleteResult {
let Some(slot) = slot_for(&request.r#ref, &request.argument.name) else {
return nothing();
};
if !gate.offers(slot.surface()) {
return nothing();
}
if !limiter.allow() {
tracing::debug!(
slot = ?slot,
"completion refused: this session is asking faster than the limit"
);
return nothing();
}
let candidates = match gather(ctx, slot).await {
Ok(candidates) => candidates,
Err(why) => {
tracing::debug!(slot = ?slot, %why, "completion found nothing to offer");
return nothing();
}
};
answer(ctx, &candidates, &request.argument.value)
}
fn answer(ctx: &ToolContext, candidates: &[Candidate], typed: &str) -> CompleteResult {
let typed = normalise(typed);
let mut matched: Vec<(u8, Kind, &str)> = candidates
.iter()
.filter_map(|candidate| {
Some((
candidate.rank(&typed)?,
candidate.kind,
candidate.value.as_str(),
))
})
.collect();
matched.sort_unstable();
let total = matched.len();
let values: Vec<String> = matched
.into_iter()
.take(CompletionInfo::MAX_VALUES)
.map(|(_, _, value)| ctx.redactor.apply(value).into_owned())
.collect();
let has_more = total > values.len();
let mut completion = CompletionInfo::new(values).unwrap_or_default();
completion.total = u32::try_from(total).ok();
completion.has_more = Some(has_more);
CompleteResult::new(completion)
}
fn nothing() -> CompleteResult {
CompleteResult::new(CompletionInfo::default())
}
async fn gather(ctx: &ToolContext, slot: Slot) -> crate::error::ToolResult<Vec<Candidate>> {
match slot {
Slot::Device => devices(ctx).await,
Slot::Peer => peers(ctx).await,
Slot::Subject => subjects(ctx).await,
}
}
async fn devices(ctx: &ToolContext) -> crate::error::ToolResult<Vec<Candidate>> {
Ok(ctx
.tailnet_devices()
.await?
.iter()
.filter(|device| !device.name.is_empty())
.map(|device| {
let mut also = vec![
device.hostname.clone(),
device.short_name().to_owned(),
device.node_id.clone(),
];
also.extend(device.addresses.iter().cloned());
Candidate::new(device.name.clone(), Kind::Device, also)
})
.collect())
}
async fn peers(ctx: &ToolContext) -> crate::error::ToolResult<Vec<Candidate>> {
let Some(status) = crate::cli::status_document(&*ctx.local).await else {
return Ok(Vec::new());
};
let Some(peers) = status["Peer"].as_object() else {
return Ok(Vec::new());
};
Ok(peers.values().filter_map(peer_candidate).collect())
}
fn peer_candidate(peer: &serde_json::Value) -> Option<Candidate> {
let dns = peer["DNSName"].as_str().unwrap_or_default();
let hostname = peer["HostName"].as_str().unwrap_or_default();
let short = dns.split('.').next().unwrap_or_default();
let value = if short.is_empty() { hostname } else { short };
if value.is_empty() {
return None;
}
let mut also = vec![hostname.to_owned(), dns.to_owned()];
if let Some(addresses) = peer["TailscaleIPs"].as_array() {
also.extend(
addresses
.iter()
.filter_map(|address| Some(address.as_str()?.to_owned())),
);
}
Some(Candidate::new(value, Kind::Device, also))
}
async fn subjects(ctx: &ToolContext) -> crate::error::ToolResult<Vec<Candidate>> {
let client = ctx.tailnet()?;
let mut found: Vec<(Kind, String)> = Vec::new();
let users = client
.get(client.tailnet_path(None, "/users"))
.send_as::<serde_json::Value>()
.await?;
if let Some(listed) = users["users"].as_array() {
found.extend(
listed
.iter()
.filter_map(|user| Some((Kind::User, user["loginName"].as_str()?.to_owned()))),
);
}
match client
.get(client.tailnet_path(None, "/acl"))
.send_as::<serde_json::Value>()
.await
{
Ok(policy) => {
if let Some(owners) = policy["tagOwners"].as_object() {
found.extend(owners.keys().map(|tag| (Kind::Tag, tag.clone())));
}
}
Err(why) => tracing::debug!(%why, "completion could not read the policy file for its tags"),
}
if let Ok(devices) = ctx.tailnet_devices().await {
for device in devices.iter() {
found.extend(device.tags.iter().map(|tag| (Kind::Tag, tag.clone())));
if !device.name.is_empty() {
found.push((Kind::Device, device.name.clone()));
}
}
}
found.sort_unstable_by(|a, b| a.1.cmp(&b.1).then(a.0.cmp(&b.0)));
found.dedup_by(|a, b| a.1 == b.1);
Ok(found
.into_iter()
.filter(|(_, subject)| !subject.is_empty())
.map(|(kind, subject)| Candidate::new(subject, kind, []))
.collect())
}
#[cfg(test)]
mod tests {
use std::time::Duration;
use super::*;
#[test]
fn a_burst_is_allowed_up_to_the_reserve_and_then_refused() {
let limiter = Limiter::new();
for spent in 0..Limiter::BURST as usize {
assert!(
limiter.allow(),
"the reserve is {} and this is request {}",
Limiter::BURST,
spent + 1
);
}
assert!(
!limiter.allow(),
"the reserve is spent, so the next one waits"
);
}
#[test]
fn the_reserve_refills() {
let limiter = Limiter::new();
while limiter.allow() {}
{
let mut bucket = limiter.bucket.lock().expect("the bucket");
bucket.last -= Duration::from_secs(1);
}
assert!(
limiter.allow(),
"a second's worth of refill is {} requests",
Limiter::PER_SECOND
);
}
#[test]
fn a_peer_answers_to_every_spelling_and_is_offered_as_the_short_one() {
let peer = serde_json::json!({
"HostName": "laptop-1",
"DNSName": "laptop.example-tailnet.ts.net.",
"TailscaleIPs": ["100.64.0.2"]
});
let candidate = peer_candidate(&peer).expect("a named peer");
assert_eq!(candidate.value, "laptop");
assert_eq!(
candidate.known_as,
vec![
"100.64.0.2",
"laptop",
"laptop-1",
"laptop.example-tailnet.ts.net"
],
"no spelling should carry the root label"
);
for typed in [
"laptop.example-tailnet.ts.net",
"laptop.example-tailnet.ts.net.",
] {
assert_eq!(
candidate.rank(&normalise(typed)),
Some(0),
"`{typed}` names this peer exactly"
);
}
}
#[test]
fn a_nameless_peer_is_not_offered() {
assert!(peer_candidate(&serde_json::json!({"TailscaleIPs": ["100.64.0.9"]})).is_none());
}
#[test]
fn a_peer_without_a_magicdns_name_is_offered_as_its_hostname() {
let peer = serde_json::json!({"HostName": "printer", "DNSName": ""});
assert_eq!(peer_candidate(&peer).expect("a peer").value, "printer");
}
#[test]
fn a_slot_is_recognised_by_its_reference_and_its_argument() {
let device = Reference::for_resource(DEVICE_TEMPLATE);
assert_eq!(slot_for(&device, "device_id"), Some(Slot::Device));
assert_eq!(slot_for(&device, "id"), None);
assert_eq!(
slot_for(&Reference::for_prompt("diagnose_connectivity"), "peer"),
Some(Slot::Peer)
);
assert_eq!(
slot_for(&Reference::for_prompt("audit_tailnet_access"), "subject"),
Some(Slot::Subject)
);
assert_eq!(
slot_for(&Reference::for_prompt("review_policy_change"), "goal"),
None
);
assert_eq!(
slot_for(
&Reference::for_resource("tailnet://device/{other}"),
"other"
),
None
);
}
}