diff --git a/devolutions-gateway/src/token.rs b/devolutions-gateway/src/token.rs index bc0b74b4d..afc2c84ab 100644 --- a/devolutions-gateway/src/token.rs +++ b/devolutions-gateway/src/token.rs @@ -55,6 +55,7 @@ pub enum ContentType { WebApp, NetScan, Enrollment, + Task, } impl FromStr for ContentType { @@ -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), }), @@ -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"), } } } @@ -125,6 +128,7 @@ pub enum AccessTokenClaims { WebApp(WebAppTokenClaims), NetScan(NetScanClaims), Enrollment(EnrollmentTokenClaims), + Task(TaskTokenClaims), } impl AccessTokenClaims { @@ -140,6 +144,7 @@ impl AccessTokenClaims { AccessTokenClaims::WebApp(_) => false, AccessTokenClaims::NetScan(_) => false, AccessTokenClaims::Enrollment(_) => false, + AccessTokenClaims::Task(_) => false, } } } @@ -511,6 +516,8 @@ pub enum AccessScope { AgentDelete, #[serde(rename = "gateway.agent.read")] AgentRead, + #[serde(rename = "gateway.tasks.read")] + TasksRead, } #[derive(Clone, Serialize, Deserialize)] @@ -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)] @@ -1008,7 +1041,8 @@ fn validate_token_impl( | ContentType::Kdc | ContentType::WebApp | ContentType::NetScan - | ContentType::Enrollment => jwt.validate::(&strict_validator)?.state.claims, + | ContentType::Enrollment + | ContentType::Task => jwt.validate::(&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 @@ -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 })?; @@ -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", @@ -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 })?; @@ -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, + } + + 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) -> Result { + 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"); + } } diff --git a/tools/tokengen/src/lib.rs b/tools/tokengen/src/lib.rs index f013f1ffd..13775fdb9 100644 --- a/tools/tokengen/src/lib.rs +++ b/tools/tokengen/src/lib.rs @@ -150,6 +150,17 @@ pub struct NetScanClaim { pub jet_gw_id: Option, } +#[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, + pub exp: i64, + pub nbf: i64, + pub jti: Uuid, +} + // --- Enums --- // #[derive(Serialize, Deserialize, Clone, Copy, Debug, PartialEq)] @@ -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 --- // @@ -285,6 +303,10 @@ pub enum SubCommandArgs { revoked_jti_list: Vec, }, NetScan {}, + Task { + jet_tk: TaskKind, + jet_aid: Option, + }, } pub fn generate_token( @@ -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); diff --git a/tools/tokengen/src/main.rs b/tools/tokengen/src/main.rs index d54a591ed..4bb3d90f3 100644 --- a/tools/tokengen/src/main.rs +++ b/tools/tokengen/src/main.rs @@ -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> { @@ -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)?; @@ -287,4 +288,10 @@ enum SignSubCommand { jti: Vec, }, NetScan {}, + Task { + #[clap(long)] + jet_tk: TaskKind, + #[clap(long)] + jet_aid: Option, + }, } diff --git a/tools/tokengen/src/server/server_impl.rs b/tools/tokengen/src/server/server_impl.rs index 6be3eb125..5e3980f9b 100644 --- a/tools/tokengen/src/server/server_impl.rs +++ b/tools/tokengen/src/server/server_impl.rs @@ -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, delegation_key_path: Option) -> Router { Router::new() @@ -23,6 +23,7 @@ pub(crate) fn create_router(provisioner_key_path: Arc, 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)) } @@ -267,6 +268,23 @@ pub(crate) async fn netscan_handler( .await } +pub(crate) async fn task_handler( + Extension(provisioner_key_path): Extension>, + Extension(delegation_key_path): Extension>, + Json(request): Json, +) -> Result, (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, delegation_key_path: Option, @@ -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, +} diff --git a/utils/dotnet/Devolutions.Gateway.Utils.Tests/JsonSerializationTests.cs b/utils/dotnet/Devolutions.Gateway.Utils.Tests/JsonSerializationTests.cs index 30f165d40..d92e3eeba 100644 --- a/utils/dotnet/Devolutions.Gateway.Utils.Tests/JsonSerializationTests.cs +++ b/utils/dotnet/Devolutions.Gateway.Utils.Tests/JsonSerializationTests.cs @@ -186,6 +186,28 @@ public void ScopeClaimsAgentRead() Assert.Equal(EXPECTED, result); } + [Fact] + public void ScopeClaimsTasksRead() + { + const string EXPECTED = """{"scope":"gateway.tasks.read","jet_gw_id":"ccbaad3f-4627-4666-8bb5-cb6a1a7db815"}"""; + + var claims = new ScopeClaims(gatewayId, AccessScope.GatewayTasksRead); + string result = JsonSerializer.Serialize(claims); + Assert.Equal(EXPECTED, result); + } + + [Fact] + public void TaskClaimsForAiLog() + { + const string EXPECTED = """{"jet_tk":"ai-log","jet_aid":"3e7c1854-f1eb-42d2-b9cb-9303036e50da","jet_gw_id":"ccbaad3f-4627-4666-8bb5-cb6a1a7db815"}"""; + + var claims = TaskClaims.ForAiLog(gatewayId, sessionId); + string result = JsonSerializer.Serialize(claims); + Assert.Equal(EXPECTED, result); + Assert.Equal("TASK", claims.GetContentType()); + Assert.Equal(600, claims.GetDefaultLifetime()); + } + [Fact] public void EnrollmentClaimsAllFields() { diff --git a/utils/dotnet/Devolutions.Gateway.Utils/src/AccessScope.cs b/utils/dotnet/Devolutions.Gateway.Utils/src/AccessScope.cs index bc953a2b1..3545de781 100644 --- a/utils/dotnet/Devolutions.Gateway.Utils/src/AccessScope.cs +++ b/utils/dotnet/Devolutions.Gateway.Utils/src/AccessScope.cs @@ -30,6 +30,7 @@ internal AccessScope(string value) public static AccessScope GatewayNetMonitorDrain = new AccessScope("gateway.net.monitor.drain"); public static AccessScope GatewayAgentDelete = new AccessScope("gateway.agent.delete"); public static AccessScope GatewayAgentRead = new AccessScope("gateway.agent.read"); + public static AccessScope GatewayTasksRead = new AccessScope("gateway.tasks.read"); public override string? ToString() { diff --git a/utils/dotnet/Devolutions.Gateway.Utils/src/TaskClaims.cs b/utils/dotnet/Devolutions.Gateway.Utils/src/TaskClaims.cs new file mode 100644 index 000000000..9d50cd4a1 --- /dev/null +++ b/utils/dotnet/Devolutions.Gateway.Utils/src/TaskClaims.cs @@ -0,0 +1,45 @@ +using System.Text.Json.Serialization; + +namespace Devolutions.Gateway.Utils; + +public class TaskClaims : IGatewayClaims +{ + [JsonPropertyName("jet_tk")] + public TaskKind TaskKind { get; set; } + + [JsonPropertyName("jet_aid")] + [JsonIgnore(Condition = JsonIgnoreCondition.WhenWritingNull)] + public Guid? SessionId { get; set; } + + [JsonPropertyName("jet_gw_id")] + public Guid ScopeGatewayId { get; set; } + + private TaskClaims(Guid scopeGatewayId, TaskKind taskKind) + { + this.ScopeGatewayId = scopeGatewayId; + this.TaskKind = taskKind; + } + + /// + /// Build the claims of a task that describes what the user did in one session and stores the result as a new log of that session. + /// + /// Target Gateway identifier. + /// Session to describe. + public static TaskClaims ForAiLog(Guid scopeGatewayId, Guid sessionId) + { + return new TaskClaims(scopeGatewayId, TaskKind.AiLog) + { + SessionId = sessionId, + }; + } + + public string GetContentType() + { + return "TASK"; + } + + public long? GetDefaultLifetime() + { + return 600; + } +} diff --git a/utils/dotnet/Devolutions.Gateway.Utils/src/TaskKind.cs b/utils/dotnet/Devolutions.Gateway.Utils/src/TaskKind.cs new file mode 100644 index 000000000..7018c00af --- /dev/null +++ b/utils/dotnet/Devolutions.Gateway.Utils/src/TaskKind.cs @@ -0,0 +1,21 @@ +using System.Text.Json.Serialization; + +namespace Devolutions.Gateway.Utils; + +[JsonConverter(typeof(TaskKindJsonConverter))] +public struct TaskKind +{ + public string Value { get; internal set; } + + internal TaskKind(string value) + { + Value = value; + } + + public static TaskKind AiLog = new TaskKind("ai-log"); + + public override string? ToString() + { + return this.Value; + } +} diff --git a/utils/dotnet/Devolutions.Gateway.Utils/src/TaskKindJsonConverter.cs b/utils/dotnet/Devolutions.Gateway.Utils/src/TaskKindJsonConverter.cs new file mode 100644 index 000000000..bba7396a5 --- /dev/null +++ b/utils/dotnet/Devolutions.Gateway.Utils/src/TaskKindJsonConverter.cs @@ -0,0 +1,17 @@ +using System.Text.Json; +using System.Text.Json.Serialization; + +namespace Devolutions.Gateway.Utils; + +public class TaskKindJsonConverter : JsonConverter +{ + public override TaskKind Read( + ref Utf8JsonReader reader, + Type typeToConvert, + JsonSerializerOptions options) => new TaskKind(reader.GetString()!); + + public override void Write( + Utf8JsonWriter writer, + TaskKind taskKind, + JsonSerializerOptions options) => writer.WriteStringValue(taskKind.ToString()); +}