Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
155 changes: 152 additions & 3 deletions devolutions-gateway/src/token.rs
Original file line number Diff line number Diff line change
Expand Up @@ -55,6 +55,7 @@ pub enum ContentType {
WebApp,
NetScan,
Enrollment,
Task,
}

impl FromStr for ContentType {
Expand All @@ -72,6 +73,7 @@ impl FromStr for ContentType {
"WEBAPP" => Ok(ContentType::WebApp),
"NETSCAN" => Ok(ContentType::NetScan),
"ENROLLMENT" => Ok(ContentType::Enrollment),
"TASK" => Ok(ContentType::Task),
unexpected => Err(BadContentType {
value: SmolStr::new(unexpected),
}),
Expand All @@ -92,6 +94,7 @@ impl fmt::Display for ContentType {
ContentType::WebApp => write!(f, "WEBAPP"),
ContentType::NetScan => write!(f, "NETSCAN"),
ContentType::Enrollment => write!(f, "ENROLLMENT"),
ContentType::Task => write!(f, "TASK"),
}
}
}
Expand Down Expand Up @@ -125,6 +128,7 @@ pub enum AccessTokenClaims {
WebApp(WebAppTokenClaims),
NetScan(NetScanClaims),
Enrollment(EnrollmentTokenClaims),
Task(TaskTokenClaims),
}

impl AccessTokenClaims {
Expand All @@ -140,6 +144,7 @@ impl AccessTokenClaims {
AccessTokenClaims::WebApp(_) => false,
AccessTokenClaims::NetScan(_) => false,
AccessTokenClaims::Enrollment(_) => false,
AccessTokenClaims::Task(_) => false,
}
}
}
Expand Down Expand Up @@ -511,6 +516,8 @@ pub enum AccessScope {
AgentDelete,
#[serde(rename = "gateway.agent.read")]
AgentRead,
#[serde(rename = "gateway.tasks.read")]
TasksRead,
}

#[derive(Clone, Serialize, Deserialize)]
Expand Down Expand Up @@ -547,6 +554,32 @@ pub struct EnrollmentTokenClaims {
pub jet_agent_name: String,
}

// ----- task claims ----- //

/// Kind of background task, with the target that this kind works on.
#[derive(Debug, Clone, PartialEq, Eq, Deserialize)]
#[serde(tag = "jet_tk")]
pub enum TaskKind {
/// Describe what the user did in one session and store the result as a new log of that session.
#[serde(rename = "ai-log")]
AiLog {
/// Association ID (= Session ID) of the session to describe.
jet_aid: Uuid,
},
}

#[derive(Debug, Clone, Deserialize)]
pub struct TaskTokenClaims {
#[serde(flatten)]
pub kind: TaskKind,

/// JWT expiration time claim.
pub exp: i64,

/// JWT "JWT ID" claim, the unique ID for this token.
pub jti: Uuid,
}

// ----- bridge claims ----- //

#[derive(Clone)]
Expand Down Expand Up @@ -1008,7 +1041,8 @@ fn validate_token_impl(
| ContentType::Kdc
| ContentType::WebApp
| ContentType::NetScan
| ContentType::Enrollment => jwt.validate::<Value>(&strict_validator)?.state.claims,
| ContentType::Enrollment
| ContentType::Task => jwt.validate::<Value>(&strict_validator)?.state.claims,
ContentType::Jrl => {
// NOTE: JRL tokens are not expected to expire.
// However, `iat` (Issued At) claim is required, and only more recent tokens will
Expand Down Expand Up @@ -1095,6 +1129,7 @@ fn validate_token_impl(
ContentType::WebApp => serde_json::from_value(claims).map(AccessTokenClaims::WebApp),
ContentType::NetScan => serde_json::from_value(claims).map(AccessTokenClaims::NetScan),
ContentType::Enrollment => serde_json::from_value(claims).map(AccessTokenClaims::Enrollment),
ContentType::Task => serde_json::from_value(claims).map(AccessTokenClaims::Task),
}
.map_err(|source| TokenError::InvalidClaimScheme { content_type, source })?;

Expand Down Expand Up @@ -1179,10 +1214,11 @@ fn validate_token_impl(
}
}

// SCOPE, NETSCAN, and JMUX tokens can never be reused.
// SCOPE, NETSCAN, JMUX, and TASK tokens can never be reused.
AccessTokenClaims::Scope(ScopeTokenClaims { jti: id, exp, .. })
| AccessTokenClaims::NetScan(NetScanClaims { jti: id, exp, .. })
| AccessTokenClaims::Jmux(JmuxTokenClaims { jti: id, exp, .. }) => match token_cache.lock().entry(id) {
| AccessTokenClaims::Jmux(JmuxTokenClaims { jti: id, exp, .. })
| AccessTokenClaims::Task(TaskTokenClaims { jti: id, exp, .. }) => match token_cache.lock().entry(id) {
Entry::Occupied(_) => {
return Err(TokenError::UnexpectedReplay {
reason: "never allowed for this use case",
Expand Down Expand Up @@ -1444,6 +1480,7 @@ pub mod unsafe_debug {
ContentType::WebApp => serde_json::from_value(claims).map(AccessTokenClaims::WebApp),
ContentType::NetScan => serde_json::from_value(claims).map(AccessTokenClaims::NetScan),
ContentType::Enrollment => serde_json::from_value(claims).map(AccessTokenClaims::Enrollment),
ContentType::Task => serde_json::from_value(claims).map(AccessTokenClaims::Task),
}
.map_err(|source| TokenError::InvalidClaimScheme { content_type, source })?;

Expand Down Expand Up @@ -1927,4 +1964,116 @@ mod tests {
assert_eq!(recording_file_type.content_type(), expected_content_type);
}
}

struct TaskTokenFixture {
provisioner_key: PrivateKey,
token_cache: TokenCache,
revocation_list: CurrentJrl,
active_recordings: Arc<ActiveRecordings>,
}

impl TaskTokenFixture {
fn new() -> Self {
let (sender, _) = crate::recording::recording_message_channel();

Self {
provisioner_key: PrivateKey::generate_ec(picky::key::EcCurve::NistP256).expect("generate EC key"),
token_cache: new_token_cache(),
revocation_list: Mutex::new(JrlTokenClaims::default()),
active_recordings: sender.active_recordings,
}
}

fn sign(&self, claims: &serde_json::Value) -> String {
picky::jose::jwt::CheckedJwtSig::new_with_cty(picky::jose::jws::JwsAlg::ES256, "TASK", claims)
.encode(&self.provisioner_key)
.expect("sign TASK token")
}

fn validate(&self, token: &str, gw_id: Option<Uuid>) -> Result<AccessTokenClaims, TokenError> {
TokenValidator::builder()
.source_ip(IpAddr::from([127, 0, 0, 1]))
.provisioner_key(&self.provisioner_key.to_public_key().expect("public key"))
.token_cache(&self.token_cache)
.revocation_list(&self.revocation_list)
.active_recordings(&self.active_recordings)
.delegation_key(None)
.subkey(None)
.gw_id(gw_id)
.disconnected_info(None)
.build()
.validate(token)
}
}

fn task_claims(extra: serde_json::Value) -> serde_json::Value {
let mut claims = serde_json::json!({
"jet_tk": "ai-log",
"jet_aid": "5e3e833f-84c7-4541-b676-acc3299e39b8",
"nbf": time::OffsetDateTime::now_utc().unix_timestamp(),
"exp": time::OffsetDateTime::now_utc().unix_timestamp() + 600,
"jti": Uuid::new_v4(),
});

claims
.as_object_mut()
.expect("object")
.extend(extra.as_object().expect("object").clone());

claims
}

#[test]
fn task_token_claims_parse() {
let fixture = TaskTokenFixture::new();
let token = fixture.sign(&task_claims(serde_json::json!({})));

let claims = fixture.validate(&token, None).expect("valid TASK token");

let AccessTokenClaims::Task(claims) = claims else {
panic!("expected TASK claims");
};
assert_eq!(
claims.kind,
TaskKind::AiLog {
jet_aid: Uuid::parse_str("5e3e833f-84c7-4541-b676-acc3299e39b8").expect("UUID"),
}
);
}

#[test]
fn task_token_is_rejected_on_second_use() {
let fixture = TaskTokenFixture::new();
let token = fixture.sign(&task_claims(serde_json::json!({})));

fixture.validate(&token, None).expect("first use");
let error = fixture.validate(&token, None).err().expect("second use is rejected");

assert!(matches!(error, TokenError::UnexpectedReplay { .. }), "{error:?}");
}

#[test]
fn unknown_task_kind_is_rejected() {
let fixture = TaskTokenFixture::new();
let token = fixture.sign(&task_claims(serde_json::json!({ "jet_tk": "monitoring" })));

let error = fixture.validate(&token, None).err().expect("unknown kind is rejected");

assert!(matches!(error, TokenError::InvalidClaimScheme { .. }), "{error:?}");
}

#[test]
fn task_token_honors_gateway_id_scope() {
let fixture = TaskTokenFixture::new();
let gw_id = Uuid::new_v4();
let token = fixture.sign(&task_claims(serde_json::json!({ "jet_gw_id": gw_id })));

let error = fixture
.validate(&token, Some(Uuid::new_v4()))
.err()
.expect("other gateway is rejected");
assert!(matches!(error, TokenError::GatewayIdScopeMismatch), "{error:?}");

fixture.validate(&token, Some(gw_id)).expect("this gateway is accepted");
}
}
33 changes: 33 additions & 0 deletions tools/tokengen/src/lib.rs
Original file line number Diff line number Diff line change
Expand Up @@ -150,6 +150,17 @@ pub struct NetScanClaim {
pub jet_gw_id: Option<Uuid>,
}

#[derive(Clone, Serialize)]
pub struct TaskClaims {
pub jet_tk: TaskKind,
pub jet_aid: Uuid,
#[serde(skip_serializing_if = "Option::is_none")]
pub jet_gw_id: Option<Uuid>,
pub exp: i64,
pub nbf: i64,
pub jti: Uuid,
}

// --- Enums --- //

#[derive(Serialize, Deserialize, Clone, Copy, Debug, PartialEq)]
Expand Down Expand Up @@ -219,9 +230,16 @@ macro_rules! impl_from_str {
};
}

#[derive(Serialize, Deserialize, Clone, Copy, Debug, PartialEq, Eq)]
#[serde(rename_all = "kebab-case")]
pub enum TaskKind {
AiLog,
}

impl_from_str!(ApplicationProtocol);
impl_from_str!(RecordingOperation);
impl_from_str!(RecordingPolicy);
impl_from_str!(TaskKind);

// --- SubCommandArgs Enum --- //

Expand Down Expand Up @@ -285,6 +303,10 @@ pub enum SubCommandArgs {
revoked_jti_list: Vec<Uuid>,
},
NetScan {},
Task {
jet_tk: TaskKind,
jet_aid: Option<Uuid>,
},
}

pub fn generate_token(
Expand Down Expand Up @@ -549,6 +571,17 @@ pub fn generate_token(
};
("NETSCAN", serde_json::to_value(claims)?)
}
SubCommandArgs::Task { jet_tk, jet_aid } => {
let claims = TaskClaims {
jet_tk,
jet_aid: jet_aid.unwrap_or_else(Uuid::new_v4),
jet_gw_id,
exp,
nbf,
jti,
};
("TASK", serde_json::to_value(claims)?)
}
};

let mut jwt_sig = CheckedJwtSig::new_with_cty(JwsAlg::RS256, cty, claims);
Expand Down
9 changes: 8 additions & 1 deletion tools/tokengen/src/main.rs
Original file line number Diff line number Diff line change
Expand Up @@ -2,7 +2,7 @@ use std::error::Error;
use std::path::{Path, PathBuf};

use clap::{Parser, Subcommand};
use tokengen::{ApplicationProtocol, RecordingOperation, SubCommandArgs, generate_token};
use tokengen::{ApplicationProtocol, RecordingOperation, SubCommandArgs, TaskKind, generate_token};
use uuid::Uuid;

fn main() -> Result<(), Box<dyn Error>> {
Expand Down Expand Up @@ -140,6 +140,7 @@ fn sign(
},
SignSubCommand::Jrl { jti } => SubCommandArgs::Jrl { revoked_jti_list: jti },
SignSubCommand::NetScan {} => SubCommandArgs::NetScan {},
SignSubCommand::Task { jet_tk, jet_aid } => SubCommandArgs::Task { jet_tk, jet_aid },
};

let validity_duration = humantime::parse_duration(validity_duration)?;
Expand Down Expand Up @@ -287,4 +288,10 @@ enum SignSubCommand {
jti: Vec<Uuid>,
},
NetScan {},
Task {
#[clap(long)]
jet_tk: TaskKind,
#[clap(long)]
jet_aid: Option<Uuid>,
},
}
28 changes: 27 additions & 1 deletion tools/tokengen/src/server/server_impl.rs
Original file line number Diff line number Diff line change
Expand Up @@ -9,7 +9,7 @@ use axum::routing::post;
use serde::{Deserialize, Serialize};
use uuid::Uuid;

use crate::{ApplicationProtocol, RecordingOperation, SubCommandArgs, generate_token};
use crate::{ApplicationProtocol, RecordingOperation, SubCommandArgs, TaskKind, generate_token};

pub(crate) fn create_router(provisioner_key_path: Arc<PathBuf>, delegation_key_path: Option<PathBuf>) -> Router {
Router::new()
Expand All @@ -23,6 +23,7 @@ pub(crate) fn create_router(provisioner_key_path: Arc<PathBuf>, delegation_key_p
.route("/kdc", post(kdc_handler))
.route("/jrl", post(jrl_handler))
.route("/netscan", post(netscan_handler))
.route("/task", post(task_handler))
.layer(Extension(provisioner_key_path))
.layer(Extension(delegation_key_path))
}
Expand Down Expand Up @@ -267,6 +268,23 @@ pub(crate) async fn netscan_handler(
.await
}

pub(crate) async fn task_handler(
Extension(provisioner_key_path): Extension<Arc<PathBuf>>,
Extension(delegation_key_path): Extension<Option<PathBuf>>,
Json(request): Json<TaskRequest>,
) -> Result<Json<TokenResponse>, (axum::http::StatusCode, String)> {
handle_subcommand(
provisioner_key_path,
delegation_key_path,
request.common,
SubCommandArgs::Task {
jet_tk: request.jet_tk,
jet_aid: request.jet_aid,
},
)
.await
}

async fn handle_subcommand(
provisioner_key_path: Arc<PathBuf>,
delegation_key_path: Option<PathBuf>,
Expand Down Expand Up @@ -387,3 +405,11 @@ pub(crate) struct NetScanRequest {
#[serde(flatten)]
common: CommonRequest,
}

#[derive(Deserialize)]
pub(crate) struct TaskRequest {
#[serde(flatten)]
common: CommonRequest,
jet_tk: TaskKind,
jet_aid: Option<Uuid>,
}
Loading
Loading