use std::sync::Arc;
use indexmap::IndexMap;
use crate::context::IExpressionContext;
use crate::exceptions::TemplateProcessingException;
use crate::expression::TemplateValue;
use crate::util::{Utf16String, to_lower_unit};
use super::ILinkBuilder;
type LinkParameters = IndexMap<Option<Utf16String>, Option<Arc<TemplateValue>>>;
type ContextPathHook = dyn Fn(
&dyn IExpressionContext,
&Utf16String,
Option<&LinkParameters>,
) -> Result<Option<Utf16String>, TemplateProcessingException>
+ Send
+ Sync;
type ProcessLinkHook = dyn Fn(
&dyn IExpressionContext,
&Utf16String,
) -> Result<Option<Utf16String>, TemplateProcessingException>
+ Send
+ Sync;
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
enum LinkType {
Absolute,
ContextRelative,
ServerRelative,
BaseRelative,
}
pub struct StandardLinkBuilder {
name: Option<Utf16String>,
order: Option<i32>,
context_path_hook: Option<Arc<ContextPathHook>>,
process_link_hook: Option<Arc<ProcessLinkHook>>,
}
impl StandardLinkBuilder {
#[must_use]
pub fn new() -> Self {
Self {
name: Some(Utf16String::from_rust_str(
"org.thymeleaf.linkbuilder.StandardLinkBuilder",
)),
order: None,
context_path_hook: None,
process_link_hook: None,
}
}
#[must_use]
pub fn with_context_path_hook<F>(mut self, hook: F) -> Self
where
F: Fn(
&dyn IExpressionContext,
&Utf16String,
Option<&IndexMap<Option<Utf16String>, Option<Arc<TemplateValue>>>>,
) -> Result<Option<Utf16String>, TemplateProcessingException>
+ Send
+ Sync
+ 'static,
{
self.context_path_hook = Some(Arc::new(hook));
self
}
#[must_use]
pub fn with_process_link_hook<F>(mut self, hook: F) -> Self
where
F: Fn(
&dyn IExpressionContext,
&Utf16String,
) -> Result<Option<Utf16String>, TemplateProcessingException>
+ Send
+ Sync
+ 'static,
{
self.process_link_hook = Some(Arc::new(hook));
self
}
#[must_use]
pub const fn get_name(&self) -> Option<&Utf16String> {
self.name.as_ref()
}
pub fn set_name(&mut self, name: Option<Utf16String>) {
self.name = name;
}
#[must_use]
pub const fn get_order(&self) -> Option<i32> {
self.order
}
pub fn set_order(&mut self, order: Option<i32>) {
self.order = order;
}
fn build_standard_link(
&self,
context: &dyn IExpressionContext,
base: Option<&Utf16String>,
parameters: Option<&LinkParameters>,
) -> Result<Option<Utf16String>, TemplateProcessingException> {
let Some(base) = base else {
return Ok(None);
};
filter_out_java_script_links(base)?;
let link_type = classify_link(base);
let mut link_parameters = parameters.filter(|value| !value.is_empty()).cloned();
let hash_position = find_last_unit(base.as_utf16(), u16::from(b'#'));
let might_have_variable_templates =
find_last_unit(base.as_utf16(), u16::from(b'{')).is_some();
let context_path = if link_type == LinkType::ContextRelative {
self.compute_context_path(context, base, parameters)?
} else {
None
};
let context_path_empty = context_path
.as_ref()
.is_none_or(|path| path.is_empty() || path.as_utf16() == [u16::from(b'/')]);
if context_path_empty
&& link_type != LinkType::ServerRelative
&& link_parameters.as_ref().is_none_or(IndexMap::is_empty)
&& hash_position.is_none()
&& !might_have_variable_templates
{
return self.process_link(context, base);
}
let mut link_base = base.as_utf16().to_vec();
let mut url_fragment = Vec::new();
if let Some(position) = hash_position.filter(|position| *position > 0) {
url_fragment.extend_from_slice(&link_base[position..]);
link_base.truncate(position);
}
if might_have_variable_templates {
replace_template_params_in_base(&mut link_base, link_parameters.as_mut());
}
if let Some(parameters) = link_parameters.as_ref().filter(|value| !value.is_empty()) {
link_base.push(if find_last_unit(&link_base, u16::from(b'?')).is_some() {
u16::from(b'&')
} else {
u16::from(b'?')
});
process_all_remaining_parameters_as_query_params(&mut link_base, parameters);
}
link_base.extend_from_slice(&url_fragment);
if link_type == LinkType::ServerRelative {
link_base.remove(0);
}
if link_type == LinkType::ContextRelative
&& !context_path_empty
&& let Some(context_path) = context_path
{
link_base.splice(0..0, context_path.as_utf16().iter().copied());
}
self.process_link(context, &Utf16String::from_utf16(link_base))
}
fn compute_context_path(
&self,
context: &dyn IExpressionContext,
base: &Utf16String,
parameters: Option<&LinkParameters>,
) -> Result<Option<Utf16String>, TemplateProcessingException> {
if let Some(hook) = &self.context_path_hook {
return hook(context, base, parameters);
}
let Some(exchange) = context.get_web_exchange() else {
return Err(TemplateProcessingException::new(Some(format!(
"Link base \"{}\" cannot be context relative (/...) unless the context used for \
executing the engine implements the org.thymeleaf.context.IWebContext interface",
base.to_string_lossy()
))));
};
Ok(exchange.get_request().get_application_path())
}
fn process_link(
&self,
context: &dyn IExpressionContext,
link: &Utf16String,
) -> Result<Option<Utf16String>, TemplateProcessingException> {
if let Some(hook) = &self.process_link_hook {
return hook(context, link);
}
Ok(context.get_web_exchange().map_or_else(
|| Some(link.clone()),
|exchange| exchange.transform_url(Some(link)),
))
}
}
impl Default for StandardLinkBuilder {
fn default() -> Self {
Self::new()
}
}
impl ILinkBuilder for StandardLinkBuilder {
fn get_name(&self) -> Option<&Utf16String> {
self.get_name()
}
fn get_order(&self) -> Option<i32> {
self.get_order()
}
fn build_link(
&self,
context: &dyn IExpressionContext,
base: Option<&Utf16String>,
parameters: Option<&LinkParameters>,
) -> Result<Option<Utf16String>, TemplateProcessingException> {
self.build_standard_link(context, base, parameters)
}
}
fn classify_link(base: &Utf16String) -> LinkType {
if is_link_base_absolute(base) {
LinkType::Absolute
} else if is_link_base_context_relative(base) {
LinkType::ContextRelative
} else if is_link_base_server_relative(base) {
LinkType::ServerRelative
} else {
LinkType::BaseRelative
}
}
fn filter_out_java_script_links(base: &Utf16String) -> Result<(), TemplateProcessingException> {
if starts_with_java_ignore_case(base.as_utf16(), b"javascript:") {
return Err(TemplateProcessingException::new(Some(
"'javascript:' is forbidden in this context. Link expressions cannot contain inlined \
JavaScript code."
.to_owned(),
)));
}
Ok(())
}
fn is_link_base_absolute(base: &Utf16String) -> bool {
let units = base.as_utf16();
if units.len() < 2 {
return false;
}
if starts_with_java_ignore_case(units, b"mailto:") {
return true;
}
if units.starts_with(&[u16::from(b'/'), u16::from(b'/')]) {
return true;
}
units
.windows(3)
.any(|window| window == [u16::from(b':'), u16::from(b'/'), u16::from(b'/')])
}
fn is_link_base_context_relative(base: &Utf16String) -> bool {
let units = base.as_utf16();
units.first() == Some(&u16::from(b'/'))
&& units.get(1).is_none_or(|unit| *unit != u16::from(b'/'))
}
fn is_link_base_server_relative(base: &Utf16String) -> bool {
base.as_utf16()
.starts_with(&[u16::from(b'~'), u16::from(b'/')])
}
fn starts_with_java_ignore_case(units: &[u16], expected: &[u8]) -> bool {
units.len() >= expected.len()
&& units
.iter()
.zip(expected)
.all(|(unit, expected)| to_lower_unit(*unit) == u16::from(*expected))
}
fn find_last_unit(units: &[u16], needle: u16) -> Option<usize> {
units.iter().rposition(|unit| *unit == needle)
}
fn replace_template_params_in_base(
link_base: &mut Vec<u16>,
parameters: Option<&mut LinkParameters>,
) {
let Some(parameters) = parameters else {
return;
};
let question_mark_position = find_last_unit(link_base, u16::from(b'?'));
let mut processed = Vec::new();
for (parameter_name, parameter_value) in parameters.iter() {
let parameter_name_text = parameter_name
.clone()
.unwrap_or_else(|| Utf16String::from_rust_str("null"));
let direct_template = surrounded_template(parameter_name_text.as_utf16(), false);
let segment_template = surrounded_template(parameter_name_text.as_utf16(), true);
let (template, escape_as_path_segment, mut start) =
if let Some(start) = find_subsequence(link_base, &direct_template, 0) {
(direct_template, false, start)
} else if let Some(start) = find_subsequence(link_base, &segment_template, 0) {
(segment_template, true, start)
} else {
continue;
};
processed.push(parameter_name.clone());
let replacement =
format_parameter_value_as_unescaped_variable_template(parameter_value.as_deref());
let replacement_len = replacement.len();
while start < link_base.len() {
let escaped = if question_mark_position.is_none_or(|question| start < question) {
if escape_as_path_segment {
escape_uri_path_segment(&replacement)
} else {
escape_uri_path(&replacement)
}
} else {
escape_uri_query_param(&replacement)
};
link_base.splice(
start..start + template.len(),
escaped.as_utf16().iter().copied(),
);
let next_start = start.saturating_add(replacement_len);
let Some(next) = find_subsequence(link_base, &template, next_start) else {
break;
};
start = next;
}
}
for parameter_name in processed {
parameters.shift_remove(¶meter_name);
}
}
fn surrounded_template(name: &[u16], segment: bool) -> Vec<u16> {
let mut template = Vec::with_capacity(name.len() + usize::from(segment) + 2);
template.push(u16::from(b'{'));
if segment {
template.push(u16::from(b'/'));
}
template.extend_from_slice(name);
template.push(u16::from(b'}'));
template
}
fn find_subsequence(haystack: &[u16], needle: &[u16], start: usize) -> Option<usize> {
if start > haystack.len() || needle.len() > haystack.len().saturating_sub(start) {
return None;
}
haystack[start..]
.windows(needle.len())
.position(|window| window == needle)
.map(|position| start + position)
}
fn format_parameter_value_as_unescaped_variable_template(
parameter_value: Option<&TemplateValue>,
) -> Utf16String {
match parameter_value {
None | Some(TemplateValue::Null) => Utf16String::from_utf16(Vec::new()),
Some(TemplateValue::List(values)) => {
let mut result = Vec::new();
for value in values.iter() {
if !result.is_empty() {
result.push(u16::from(b','));
}
if !matches!(value.as_ref(), TemplateValue::Null) {
result.extend_from_slice(value_string(value).as_utf16());
}
}
Utf16String::from_utf16(result)
}
Some(value) => value_string(value),
}
}
fn process_all_remaining_parameters_as_query_params(
result: &mut Vec<u16>,
parameters: &LinkParameters,
) {
let mut parameter_index = 0usize;
for (parameter_name, value) in parameters {
let parameter_name = parameter_name
.clone()
.unwrap_or_else(|| Utf16String::from_rust_str("null"));
match value.as_deref() {
None | Some(TemplateValue::Null) => {
if parameter_index > 0 {
result.push(u16::from(b'&'));
}
append_utf16_string(result, &escape_uri_query_param(¶meter_name));
parameter_index += 1;
continue;
}
Some(TemplateValue::List(values)) => {
for (value_index, value) in values.iter().enumerate() {
if parameter_index > 0 || value_index > 0 {
result.push(u16::from(b'&'));
}
append_utf16_string(result, &escape_uri_query_param(¶meter_name));
if !matches!(value.as_ref(), TemplateValue::Null) {
result.push(u16::from(b'='));
append_utf16_string(result, &escape_uri_query_param(&value_string(value)));
}
}
}
Some(value) => {
if parameter_index > 0 {
result.push(u16::from(b'&'));
}
append_utf16_string(result, &escape_uri_query_param(¶meter_name));
result.push(u16::from(b'='));
append_utf16_string(result, &escape_uri_query_param(&value_string(value)));
}
}
parameter_index += 1;
}
}
fn value_string(value: &TemplateValue) -> Utf16String {
value
.to_utf16_string()
.unwrap_or_else(|| Utf16String::from_rust_str("null"))
}
fn append_utf16_string(result: &mut Vec<u16>, value: &Utf16String) {
result.extend_from_slice(value.as_utf16());
}
fn escape_uri_path(value: &Utf16String) -> Utf16String {
percent_escape(value, |byte| is_pchar(byte) || byte == b'/')
}
fn escape_uri_path_segment(value: &Utf16String) -> Utf16String {
percent_escape(value, is_pchar)
}
fn escape_uri_query_param(value: &Utf16String) -> Utf16String {
percent_escape(value, |byte| {
!matches!(byte, b'=' | b'&' | b'+' | b'#')
&& (is_pchar(byte) || matches!(byte, b'/' | b'?'))
})
}
fn is_pchar(byte: u8) -> bool {
byte.is_ascii_alphanumeric()
|| matches!(
byte,
b'-' | b'.'
| b'_'
| b'~'
| b'!'
| b'$'
| b'&'
| b'\''
| b'('
| b')'
| b'*'
| b'+'
| b','
| b';'
| b'='
| b':'
| b'@'
)
}
fn percent_escape(value: &Utf16String, allowed: impl Fn(u8) -> bool) -> Utf16String {
let units = value.as_utf16();
let mut result = Vec::with_capacity(units.len());
let mut index = 0usize;
while index < units.len() {
let unit = units[index];
if unit <= 0x7f {
let byte = unit as u8;
if allowed(byte) {
result.push(unit);
} else {
append_percent_byte(&mut result, byte);
}
index += 1;
continue;
}
if (0xd800..=0xdbff).contains(&unit)
&& let Some(low) = units.get(index + 1).copied()
&& (0xdc00..=0xdfff).contains(&low)
{
let scalar = 0x1_0000 + ((u32::from(unit) - 0xd800) << 10) + (u32::from(low) - 0xdc00);
let character = char::from_u32(scalar).expect("valid surrogate pair");
let mut buffer = [0u8; 4];
for byte in character.encode_utf8(&mut buffer).as_bytes() {
append_percent_byte(&mut result, *byte);
}
index += 2;
continue;
}
if (0xd800..=0xdfff).contains(&unit) {
append_percent_byte(&mut result, b'?');
index += 1;
continue;
}
let character = char::from_u32(u32::from(unit)).expect("non-surrogate BMP unit");
let mut buffer = [0u8; 3];
for byte in character.encode_utf8(&mut buffer).as_bytes() {
append_percent_byte(&mut result, *byte);
}
index += 1;
}
Utf16String::from_utf16(result)
}
fn append_percent_byte(result: &mut Vec<u16>, byte: u8) {
const HEX: &[u8; 16] = b"0123456789ABCDEF";
result.push(u16::from(b'%'));
result.push(u16::from(HEX[usize::from(byte >> 4)]));
result.push(u16::from(HEX[usize::from(byte & 0x0f)]));
}