use rmcp::schemars::JsonSchema;
use serde::{Deserialize, Serialize};
use serde_json::Value;
use tailscale_rest::models::policy::PREVIEW_SUBJECTS;
use crate::context::ToolContext;
use crate::error::{ToolError, ToolResult};
use crate::tools::common::{one_of, report};
crate::tools! {
tailnet_policy_get => PolicyGetParams, policy_get,
toolset: TailnetPolicy, tier: Read, idempotent: true;
tailnet_policy_set => PolicySetParams, policy_set,
toolset: TailnetPolicy, tier: Destructive;
tailnet_policy_preview => PolicyPreviewParams, policy_preview,
toolset: TailnetPolicy, tier: Read, idempotent: true;
tailnet_policy_validate => PolicyValidateParams, policy_validate,
toolset: TailnetPolicy, tier: Read, idempotent: true;
}
const FORMATS: &[&str] = &["hujson", "json"];
const HUJSON: &str = "application/hujson";
const JSON: &str = "application/json";
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
enum Format {
HuJson,
Json,
}
impl Format {
fn parse(given: Option<&str>) -> ToolResult<Self> {
match given {
None => Ok(Self::HuJson),
Some(value) => match one_of("format", value, FORMATS)?.as_str() {
"json" => Ok(Self::Json),
_ => Ok(Self::HuJson),
},
}
}
fn accept(self) -> &'static str {
match self {
Self::HuJson => HUJSON,
Self::Json => JSON,
}
}
fn as_str(self) -> &'static str {
match self {
Self::HuJson => "hujson",
Self::Json => "json",
}
}
fn read(self, text: String) -> ToolResult<Value> {
match self {
Self::HuJson => Ok(Value::String(text)),
Self::Json => serde_json::from_str(&text).map_err(|source| {
ToolError::new(
crate::error::ErrorCode::ApiError,
format!("the policy was asked for as JSON and did not parse as JSON: {source}"),
)
}),
}
}
}
#[derive(Debug, Deserialize, JsonSchema)]
pub struct PolicyGetParams {
#[serde(default)]
pub format: Option<String>,
#[serde(default)]
pub details: Option<bool>,
}
#[derive(Debug, Serialize)]
struct PolicyDocument {
#[serde(skip_serializing_if = "Option::is_none")]
etag: Option<String>,
format: &'static str,
policy: Value,
}
#[derive(Debug, Serialize)]
struct PolicyReport {
#[serde(skip_serializing_if = "Option::is_none")]
etag: Option<String>,
details: Value,
}
async fn policy_get(ctx: &ToolContext, params: PolicyGetParams) -> ToolResult<Value> {
let client = ctx.tailnet()?;
let acl = client.tailnet_path(None, "/acl");
if params.details.unwrap_or(false) {
if params.format.is_some() {
return Err(ToolError::invalid_args(
"`details` and `format` cannot be given together: the detailed report is always \
JSON, and the control plane refuses an `Accept` alongside `details`",
));
}
let answer = client
.get(acl)
.query("details", true)
.send_answer::<Value>()
.await?;
return report(PolicyReport {
etag: answer.etag,
details: answer.value,
});
}
let format = Format::parse(params.format.as_deref())?;
let body = client
.get(acl)
.header("Accept", format.accept())
.send_text()
.await?;
report(PolicyDocument {
etag: body.etag,
format: format.as_str(),
policy: format.read(body.text)?,
})
}
#[derive(Debug, Deserialize, JsonSchema)]
pub struct PolicySetParams {
pub policy: Value,
#[serde(default)]
pub etag: Option<String>,
#[serde(default)]
pub over_default: Option<bool>,
}
fn if_match(etag: Option<&str>, over_default: bool) -> ToolResult<String> {
let etag = etag.map(str::trim).filter(|e| !e.is_empty());
match (etag, over_default) {
(Some(_), true) => Err(ToolError::invalid_args(
"`etag` and `over_default` say different things: one writes over the version you \
read, the other only over the untouched default. Give one.",
)),
(Some(etag), false) => Ok(quoted(etag)),
(None, true) => Ok(quoted("ts-default")),
(None, false) => Err(ToolError::invalid_args(
"a policy write needs `etag` — the version `tailnet_policy_get` answered with — or \
`over_default: true` to write only over a tailnet's untouched default policy",
)
.with_hint(
"Without one of these the write would overwrite whatever is there, including a \
change somebody else made since you last read the policy.",
)),
}
}
fn quoted(value: &str) -> String {
if value.starts_with('"') && value.ends_with('"') && value.len() >= 2 {
return value.to_owned();
}
format!("\"{value}\"")
}
async fn policy_set(ctx: &ToolContext, params: PolicySetParams) -> ToolResult<Value> {
let client = ctx.tailnet()?;
let guard = if_match(params.etag.as_deref(), params.over_default.unwrap_or(false))?;
let request = with_policy(
client
.post(client.tailnet_path(None, "/acl"))
.header("If-Match", guard),
¶ms.policy,
)?;
let body = request.send_text().await?;
report(PolicyDocument {
etag: body.etag,
format: Format::HuJson.as_str(),
policy: Value::String(body.text),
})
}
#[derive(Debug, Deserialize, JsonSchema)]
pub struct PolicyPreviewParams {
pub policy: Value,
pub subject_type: String,
pub preview_for: String,
}
async fn policy_preview(ctx: &ToolContext, params: PolicyPreviewParams) -> ToolResult<Value> {
let client = ctx.tailnet()?;
let subject_type = one_of("subject_type", ¶ms.subject_type, PREVIEW_SUBJECTS)?;
let request = client
.post(client.tailnet_path(None, "/acl/preview"))
.query("type", &subject_type)
.query("previewFor", ¶ms.preview_for);
Ok(with_policy(request, ¶ms.policy)?
.send_as::<Value>()
.await?)
}
fn with_policy<'a>(
request: tailscale_rest::RequestBuilder<'a>,
policy: &Value,
) -> ToolResult<tailscale_rest::RequestBuilder<'a>> {
match policy {
Value::String(text) => Ok(request.text(HUJSON, text.clone())),
object @ Value::Object(_) => Ok(request.json(object)),
_ => Err(ToolError::invalid_args(
"`policy` is a HuJSON document as a string or a policy object; it is neither",
)),
}
}
#[derive(Debug, Deserialize, JsonSchema)]
pub struct PolicyValidateParams {
#[serde(default)]
pub tests: Option<Vec<Value>>,
#[serde(default)]
pub policy: Option<Value>,
}
#[derive(Debug, Serialize)]
struct Passed {
passed: bool,
checked: &'static str,
}
async fn policy_validate(ctx: &ToolContext, params: PolicyValidateParams) -> ToolResult<Value> {
let client = ctx.tailnet()?;
let (request, checked) = match (params.tests, params.policy) {
(Some(_), Some(_)) => {
return Err(ToolError::invalid_args(
"`tests` and `policy` are the endpoint's two modes and only one can be sent: \
`tests` runs against the policy in force, `policy` checks a document that is \
not saved",
));
}
(None, None) => {
return Err(ToolError::invalid_args(
"give `tests` to run access tests against the current policy, or `policy` to \
check a document",
));
}
(Some(tests), None) => (
client
.post(client.tailnet_path(None, "/acl/validate"))
.json(&tests),
"the tests against the current policy",
),
(None, Some(policy)) => (
with_policy(
client.post(client.tailnet_path(None, "/acl/validate")),
&policy,
)?,
"the policy given",
),
};
let answer = request.send_as::<Value>().await?;
match answer {
Value::Null => report(Passed {
passed: true,
checked,
}),
found => Ok(found),
}
}
#[cfg(test)]
mod tests {
use super::*;
use serde_json::json;
#[test]
fn a_guard_is_quoted_however_it_arrives() {
assert_eq!(
if_match(Some("\"e0b2816b418\""), false).expect("an etag"),
"\"e0b2816b418\""
);
assert_eq!(
if_match(Some("e0b2816b418"), false).expect("an etag"),
"\"e0b2816b418\""
);
assert_eq!(if_match(None, true).expect("the default"), "\"ts-default\"");
}
#[test]
fn the_two_guards_together_are_refused_rather_than_ranked() {
assert!(if_match(Some("\"abc\""), true).is_err());
assert!(if_match(Some(" "), false).is_err());
assert_eq!(
if_match(Some(" "), true).expect("the default"),
"\"ts-default\""
);
}
#[test]
fn a_hujson_document_is_never_parsed_and_a_json_one_always_is() {
let commented = "{\n // kept\n \"acls\": [],\n}".to_owned();
assert_eq!(
Format::HuJson.read(commented.clone()).expect("text"),
json!(commented),
"parsing it would fail, and reformatting it would lose the comment"
);
assert_eq!(
Format::Json
.read("{\"acls\": []}".to_owned())
.expect("json"),
json!({"acls": []})
);
assert!(
Format::Json.read(commented).is_err(),
"JSON was asked for and something else arrived"
);
}
}