Vendor dependencies

This commit is contained in:
2026-08-01 16:11:49 +03:00
parent 7f139a0241
commit 6b5e7f0f8b
29706 changed files with 9575646 additions and 0 deletions
@@ -0,0 +1,347 @@
use insta::{assert_debug_snapshot, assert_snapshot};
use std::collections::HashMap;
use std::path::Path; // For creating regex filters
// Import only the essential functions from build/embedded_assets.rs
// Use a module declaration with the `#[path]` attribute to specify the file path
#[path = "../../build/embedded_assets.rs"]
#[allow(clippy::module_inception)]
mod embedded_assets;
// Export only the functions we're actually testing
pub use embedded_assets::{
build_static_assets, collect_all_files, discover_all_directories, find_app_directory,
generate_asset_code, generate_empty_asset_files,
};
/// Creates a test file structure with common assets for testing.
fn create_test_assets() -> tree_fs::Tree {
tree_fs::TreeBuilder::default()
.drop(true)
.add_file("assets/css/style.css", "body { color: blue; }")
.add_file("assets/js/app.js", "console.log('Hello Loco');")
.add_file("assets/views/index.html", "<h1>Hello World</h1>")
.add_directory("generated")
.create()
.unwrap()
}
/// Creates insta settings with common filters.
fn create_insta_settings(root_path: &Path) -> insta::Settings {
let mut settings = insta::Settings::clone_current();
settings.add_filter(root_path.to_str().unwrap(), "[TEST_ROOT]");
settings.add_filter("\\\\\\\\", "/");
settings
}
#[test]
fn test_generate_empty_asset_files() {
// Create a temporary test environment
let tree_fs = tree_fs::TreeBuilder::default()
.drop(true)
.add_directory("generated")
.create()
.unwrap();
let output_path = tree_fs.root.join("generated");
generate_empty_asset_files(&output_path).unwrap();
// Verify files exist
let static_file_path = output_path.join("static_assets.rs");
let templates_file_path = output_path.join("view_templates.rs");
assert!(
static_file_path.exists(),
"Static assets file should be created"
);
assert!(
templates_file_path.exists(),
"Templates file should be created"
);
// Use snapshots to verify file contents
let static_file_content = std::fs::read_to_string(static_file_path).unwrap();
let templates_file_content = std::fs::read_to_string(templates_file_path).unwrap();
assert_snapshot!("empty_static_assets_rs", static_file_content);
assert_snapshot!("empty_templates_rs", templates_file_content);
}
#[test]
fn test_generate_asset_code() {
let tree_fs = create_test_assets();
let root_path = &tree_fs.root;
let output_path = root_path.join("generated");
// Create file mapping
let mut all_files = HashMap::new();
all_files.insert(
root_path
.join("assets/css/style.css")
.to_str()
.unwrap()
.to_string(),
"/css/style.css".to_string(),
);
all_files.insert(
root_path
.join("assets/js/app.js")
.to_str()
.unwrap()
.to_string(),
"/js/app.js".to_string(),
);
all_files.insert(
root_path
.join("assets/views/index.html")
.to_str()
.unwrap()
.to_string(),
"index.html".to_string(),
);
generate_asset_code(&all_files, &output_path).unwrap();
// Verify files exist
let static_assets_path = output_path.join("static_assets.rs");
let view_templates_path = output_path.join("view_templates.rs");
assert!(static_assets_path.exists());
assert!(view_templates_path.exists());
// Snapshot file contents
let static_content = std::fs::read_to_string(static_assets_path).unwrap();
let template_content = std::fs::read_to_string(view_templates_path).unwrap();
let settings = create_insta_settings(root_path);
settings.bind(|| {
assert_snapshot!("static_assets_rs", static_content);
assert_snapshot!("view_templates_rs", template_content);
});
}
#[test]
fn test_discover_all_directories() {
let tree = tree_fs::TreeBuilder::default()
.drop(true)
.add_directory("my_assets/css")
.add_file("my_assets/css/style.css", "/* some css */")
.add_directory("my_assets/js/vendor")
.add_directory("my_assets/images")
.add_directory("my_assets/views/user")
.add_file("my_assets/root_file.txt", "I am a file")
.add_directory("my_assets/empty_dir")
.create()
.unwrap();
let assets_root_path = tree.root.join("my_assets");
let discovered_dirs = discover_all_directories(&assets_root_path);
// Use insta settings for path normalization
let settings = create_insta_settings(&tree.root);
settings.bind(|| {
assert_debug_snapshot!("discovered_directories", discovered_dirs);
});
// Test edge cases
let non_existent_path = tree.root.join("non_existent_assets");
let discovered_for_non_existent = discover_all_directories(&non_existent_path);
assert!(
discovered_for_non_existent.is_empty(),
"Should return empty for a non-existent path"
);
// Test with an empty root directory
let empty_root_path = tree.root.join("actually_empty_assets");
std::fs::create_dir(&empty_root_path).unwrap();
let discovered_for_empty = discover_all_directories(&empty_root_path);
assert_eq!(
discovered_for_empty.len(),
1,
"Should find only the root for an empty existing directory"
);
}
#[test]
fn test_find_app_directory() {
// Case 1: Standard project structure
let tree_target = tree_fs::TreeBuilder::default()
.drop(true)
.add_file("my_project/Cargo.toml", "[package]\nname = \"my_project\"")
.add_directory("my_project/src")
.add_directory("my_project/target/debug/deps")
.create()
.unwrap();
let expected_project_root = tree_target.root.join("my_project");
let out_dir_in_target = expected_project_root.join("target/debug/deps");
let app_dir = find_app_directory(&out_dir_in_target);
assert!(
app_dir.is_some(),
"Should find app directory in standard structure"
);
if let Some(found_dir) = app_dir {
assert_eq!(
found_dir, expected_project_root,
"Should find correct project root"
);
assert!(
found_dir.join("Cargo.toml").exists(),
"Located app_dir should contain Cargo.toml"
);
}
// Case 2: Path not within a 'target' directory
let tree_no_target = tree_fs::TreeBuilder::default()
.drop(true)
.add_directory("some_other_place/src")
.create()
.unwrap();
let path_not_in_target = tree_no_target.root.join("some_other_place/src");
let app_dir = find_app_directory(&path_not_in_target);
assert!(app_dir.is_some(), "Should return fallback directory");
if let Some(found_dir) = app_dir {
assert!(found_dir.exists(), "Fallback directory should exist");
}
}
#[test]
fn test_build_static_assets() {
// Create test environment
let tree_fs = tree_fs::TreeBuilder::default()
.drop(true)
.add_file("assets/css/style.css", "body { color: blue; }")
.add_file("assets/js/app.js", "console.log('Hello Loco');")
.add_file("assets/views/index.html", "<h1>Hello World</h1>")
.add_directory("target/debug/build/embedded_code")
.create()
.unwrap();
let root_path = &tree_fs.root;
let out_dir = root_path.join("target/debug/build/embedded_code");
// Call function being tested
build_static_assets(&out_dir);
// Verify files were generated
let generated_path = out_dir.join("generated_code");
let static_assets_path = generated_path.join("static_assets.rs");
let view_templates_path = generated_path.join("view_templates.rs");
assert!(generated_path.exists(), "Generated directory should exist");
assert!(
static_assets_path.exists(),
"Static assets file should exist"
);
assert!(
view_templates_path.exists(),
"View templates file should exist"
);
// Snapshot generated files
let static_content = std::fs::read_to_string(static_assets_path).unwrap();
let template_content = std::fs::read_to_string(view_templates_path).unwrap();
let settings = create_insta_settings(root_path);
settings.bind(|| {
assert_snapshot!("build_static_assets_static", static_content);
assert_snapshot!("build_static_assets_templates", template_content);
});
}
#[test]
fn test_collect_all_files() {
let tree_fs = create_test_assets();
let root_path = &tree_fs.root;
let assets_dir = root_path.join("assets");
// Test collection from css directory
let mut all_files = HashMap::new();
collect_all_files(&assets_dir.join("css"), &assets_dir, &mut all_files);
// Convert to sorted vector for consistent order
let mut file_mappings: Vec<(String, String)> = all_files
.iter()
.map(|(path, key)| (path.clone(), key.clone()))
.collect();
file_mappings.sort();
// Use insta settings for path normalization
let settings = create_insta_settings(root_path);
settings.bind(|| {
assert_debug_snapshot!("collected_css_files", file_mappings);
});
// Test collection from all directories
let mut all_files = HashMap::new();
for dir in discover_all_directories(&assets_dir) {
collect_all_files(&dir, &assets_dir, &mut all_files);
}
let mut file_mappings: Vec<(String, String)> = all_files
.iter()
.map(|(path, key)| (path.clone(), key.clone()))
.collect();
file_mappings.sort();
settings.bind(|| {
assert_debug_snapshot!("collected_all_files", file_mappings);
});
}
#[test]
fn test_template_inheritance() {
// Create test environment with complex template inheritance (4 levels)
let tree_fs = tree_fs::TreeBuilder::default()
.drop(true)
// Level 1 (base)
.add_file(
"assets/views/base.html",
"<!DOCTYPE html><html><head><title>{% block meta_title %}Base{% endblock %}</title>{% block head %}{% endblock %}</head><body>{% block body %}{% endblock %}</body></html>"
)
// Level 2 (extends base)
.add_file(
"assets/views/layouts/app.html",
"{% extends \"base.html\" %}{% block head %}<link rel=\"stylesheet\" href=\"/app.css\">{% endblock %}{% block body %}<nav>{% block nav %}{% endblock %}</nav><main>{% block content %}{% endblock %}</main>{% endblock %}"
)
// Level 3 (extends app)
.add_file(
"assets/views/layouts/authenticated.html",
"{% extends \"layouts/app.html\" %}{% block nav %}<div class=\"user-nav\">{% block user_nav %}{% endblock %}</div>{% endblock %}"
)
// Level 4 (extends authenticated)
.add_file(
"assets/views/dashboard/index.html",
"{% extends \"layouts/authenticated.html\" %}{% block meta_title %}Dashboard{% endblock %}{% block user_nav %}<a href=\"/profile\">Profile</a>{% endblock %}{% block content %}<h1>Dashboard</h1>{% endblock %}"
)
// Another Level 4 template to test multiple children
.add_file(
"assets/views/dashboard/settings.html",
"{% extends \"layouts/authenticated.html\" %}{% block meta_title %}Settings{% endblock %}{% block user_nav %}<a href=\"/profile\">Profile</a>{% endblock %}{% block content %}<h1>Settings</h1>{% endblock %}"
)
// Independent template with no inheritance
.add_file(
"assets/views/error.html",
"<h1>Error</h1>"
)
.add_directory("target/debug/build/embedded_code")
.create()
.unwrap();
let root_path = &tree_fs.root;
let out_dir = root_path.join("target/debug/build/embedded_code");
// Call function being tested
build_static_assets(&out_dir);
// Read and snapshot the generated code
let generated_path = out_dir.join("generated_code");
let view_templates_path = generated_path.join("view_templates.rs");
let template_content = std::fs::read_to_string(view_templates_path).unwrap();
let settings = create_insta_settings(root_path);
settings.bind(|| {
assert_snapshot!("complex_template_inheritance", template_content);
});
}
@@ -0,0 +1 @@
mod embedded_assets;
@@ -0,0 +1,12 @@
---
source: tests/build_scripts/embedded_assets.rs
expression: static_content
snapshot_kind: text
---
#[must_use]
pub fn get_embedded_static_assets() -> std::collections::HashMap<String, &'static [u8]> {
let mut assets = std::collections::HashMap::new();
assets.insert("/css/style.css".to_string(), include_bytes!("[TEST_ROOT]/assets/css/style.css") as &[u8]);
assets.insert("/js/app.js".to_string(), include_bytes!("[TEST_ROOT]/assets/js/app.js") as &[u8]);
assets
}
@@ -0,0 +1,14 @@
---
source: tests/build_scripts/embedded_assets.rs
expression: template_content
snapshot_kind: text
---
/// Returns a BTreeMap of templates in dependency order (parents before children)
#[must_use]
pub fn get_embedded_templates() -> std::collections::BTreeMap<String, &'static str> {
let mut templates = std::collections::BTreeMap::new();
// Base template with no parent
templates.insert("index.html".to_string(), include_str!("[TEST_ROOT]/assets/views/index.html"));
templates
}
@@ -0,0 +1,19 @@
---
source: tests/build_scripts/embedded_assets.rs
expression: file_mappings
snapshot_kind: text
---
[
(
"[TEST_ROOT]/assets/css/style.css",
"/css/style.css",
),
(
"[TEST_ROOT]/assets/js/app.js",
"/js/app.js",
),
(
"[TEST_ROOT]/assets/views/index.html",
"index.html",
),
]
@@ -0,0 +1,11 @@
---
source: tests/build_scripts/embedded_assets.rs
expression: file_mappings
snapshot_kind: text
---
[
(
"[TEST_ROOT]/assets/css/style.css",
"/css/style.css",
),
]
@@ -0,0 +1,24 @@
---
source: tests/build_scripts/embedded_assets.rs
expression: template_content
snapshot_kind: text
---
/// Returns a BTreeMap of templates in dependency order (parents before children)
#[must_use]
pub fn get_embedded_templates() -> std::collections::BTreeMap<String, &'static str> {
let mut templates = std::collections::BTreeMap::new();
// Base template with no parent
templates.insert("base.html".to_string(), include_str!("[TEST_ROOT]/assets/views/base.html"));
// Base template with no parent
templates.insert("error.html".to_string(), include_str!("[TEST_ROOT]/assets/views/error.html"));
// Template that extends base.html
templates.insert("layouts/app.html".to_string(), include_str!("[TEST_ROOT]/assets/views/layouts/app.html"));
// Template that extends layouts/app.html
templates.insert("layouts/authenticated.html".to_string(), include_str!("[TEST_ROOT]/assets/views/layouts/authenticated.html"));
// Template that extends layouts/authenticated.html
templates.insert("dashboard/index.html".to_string(), include_str!("[TEST_ROOT]/assets/views/dashboard/index.html"));
// Template that extends layouts/authenticated.html
templates.insert("dashboard/settings.html".to_string(), include_str!("[TEST_ROOT]/assets/views/dashboard/settings.html"));
templates
}
@@ -0,0 +1,15 @@
---
source: tests/build_scripts/embedded_assets.rs
expression: discovered_dirs
snapshot_kind: text
---
[
"[TEST_ROOT]/my_assets",
"[TEST_ROOT]/my_assets/css",
"[TEST_ROOT]/my_assets/empty_dir",
"[TEST_ROOT]/my_assets/images",
"[TEST_ROOT]/my_assets/js",
"[TEST_ROOT]/my_assets/js/vendor",
"[TEST_ROOT]/my_assets/views",
"[TEST_ROOT]/my_assets/views/user",
]
@@ -0,0 +1,10 @@
---
source: tests/build_scripts/embedded_assets.rs
expression: static_file_content
snapshot_kind: text
---
#[must_use]
pub fn get_embedded_static_assets() -> std::collections::HashMap<String, &'static [u8]> {
// No assets found
std::collections::HashMap::new()
}
@@ -0,0 +1,10 @@
---
source: tests/build_scripts/embedded_assets.rs
expression: templates_file_content
snapshot_kind: text
---
#[must_use]
pub fn get_embedded_templates() -> std::collections::HashMap<String, &'static str> {
// No templates found
std::collections::HashMap::new()
}
@@ -0,0 +1,12 @@
---
source: tests/build_scripts/embedded_assets.rs
expression: static_content
snapshot_kind: text
---
#[must_use]
pub fn get_embedded_static_assets() -> std::collections::HashMap<String, &'static [u8]> {
let mut assets = std::collections::HashMap::new();
assets.insert("/css/style.css".to_string(), include_bytes!("[TEST_ROOT]/assets/css/style.css") as &[u8]);
assets.insert("/js/app.js".to_string(), include_bytes!("[TEST_ROOT]/assets/js/app.js") as &[u8]);
assets
}
@@ -0,0 +1,15 @@
---
source: tests/build_scripts/embedded_assets.rs
expression: template_content
snapshot_kind: text
---
#[must_use]
pub fn get_embedded_templates() -> std::collections::HashMap<String, &'static str> {
let mut templates = std::collections::HashMap::new();
// Debug log of template keys for inheritance:
// Template key: "base.html"
// Template key: "posts/list.html"
templates.insert("base.html".to_string(), include_str!("[TEST_ROOT]/assets/views/base.html"));
templates.insert("posts/list.html".to_string(), include_str!("[TEST_ROOT]/assets/views/posts/list.html"));
templates
}
@@ -0,0 +1,14 @@
---
source: tests/build_scripts/embedded_assets.rs
expression: template_content
snapshot_kind: text
---
/// Returns a BTreeMap of templates in dependency order (parents before children)
#[must_use]
pub fn get_embedded_templates() -> std::collections::BTreeMap<String, &'static str> {
let mut templates = std::collections::BTreeMap::new();
// Base template with no parent
templates.insert("index.html".to_string(), include_str!("[TEST_ROOT]/assets/views/index.html"));
templates
}
@@ -0,0 +1,149 @@
use loco_rs::{controller::extractor::auth, prelude::*, tests_cfg};
use serde::{Deserialize, Serialize};
use loco_rs::model::{Authenticable, ModelError};
use crate::infra_cfg;
#[derive(Debug, Deserialize, Serialize)]
pub struct TestUserResponse {
pub pid: String,
pub user_id: i32,
pub user_email: String,
}
// Mock user struct for testing ApiToken extractor
#[derive(Debug, Clone)]
struct TestUser {
id: i32,
email: String,
}
#[async_trait::async_trait]
impl Authenticable for TestUser {
async fn find_by_claims_key(
_db: &sea_orm::DatabaseConnection,
pid: &str,
) -> Result<Self, ModelError> {
// Simple mock: return user if pid matches, otherwise not found
if pid == "test_pid_123" {
Ok(Self {
id: 1,
email: "test@example.com".to_string(),
})
} else {
Err(ModelError::EntityNotFound)
}
}
async fn find_by_api_key(
_db: &sea_orm::DatabaseConnection,
api_key: &str,
) -> Result<Self, ModelError> {
// Simple mock: return user if api_key matches, otherwise not found
if api_key == "test_api_key_123" {
Ok(Self {
id: 1,
email: "test@example.com".to_string(),
})
} else {
Err(ModelError::EntityNotFound)
}
}
}
// Test handler for ApiToken extractor
async fn api_token_handler(auth: auth::ApiToken<TestUser>) -> Result<Response> {
format::json(TestUserResponse {
pid: String::new(), // API tokens don't have PIDs
user_id: auth.user.id,
user_email: auth.user.email,
})
}
// Test ApiToken extractor with valid API key
#[tokio::test]
async fn can_extract_api_token_valid() {
let ctx = tests_cfg::app::get_app_context().await;
let port = get_available_port().await;
let handle =
infra_cfg::server::start_with_route(ctx, "/", get(api_token_handler), Some(port)).await;
let client = reqwest::Client::new();
let res = client
.get(get_base_url_port(port))
.header("Authorization", "Bearer test_api_key_123")
.send()
.await
.expect("Valid response");
assert_eq!(res.status(), 200);
let body: TestUserResponse = res.json().await.expect("Valid JSON response");
assert_eq!(body.pid, ""); // API tokens don't have PIDs
assert_eq!(body.user_id, 1);
assert_eq!(body.user_email, "test@example.com");
handle.abort();
}
// Test ApiToken extractor with invalid API key
#[tokio::test]
async fn can_handle_api_token_invalid() {
let ctx = tests_cfg::app::get_app_context().await;
let port = get_available_port().await;
let handle =
infra_cfg::server::start_with_route(ctx, "/", get(api_token_handler), Some(port)).await;
let client = reqwest::Client::new();
let res = client
.get(get_base_url_port(port))
.header("Authorization", "Bearer invalid_api_key")
.send()
.await
.expect("Valid response");
assert_eq!(res.status(), 401);
handle.abort();
}
// Test ApiToken extractor with missing Authorization header
#[tokio::test]
async fn can_handle_api_token_missing() {
let ctx = tests_cfg::app::get_app_context().await;
let port = get_available_port().await;
let handle =
infra_cfg::server::start_with_route(ctx, "/", get(api_token_handler), Some(port)).await;
let client = reqwest::Client::new();
let res = client
.get(get_base_url_port(port))
.send()
.await
.expect("Valid response");
assert_eq!(res.status(), 401);
handle.abort();
}
// Test response serialization
#[tokio::test]
async fn test_user_response_serialization() {
let response = TestUserResponse {
pid: "test_pid".to_string(),
user_id: 1,
user_email: "test@example.com".to_string(),
};
let json = serde_json::to_string(&response).expect("Should serialize");
let deserialized: TestUserResponse = serde_json::from_str(&json).expect("Should deserialize");
assert_eq!(response.pid, deserialized.pid);
assert_eq!(response.user_id, deserialized.user_id);
assert_eq!(response.user_email, deserialized.user_email);
}
File diff suppressed because it is too large Load Diff
@@ -0,0 +1,206 @@
use loco_rs::{controller::extractor::auth, prelude::*, tests_cfg};
use serde::{Deserialize, Serialize};
use loco_rs::model::{Authenticable, ModelError};
use crate::infra_cfg;
#[derive(Debug, Deserialize, Serialize)]
pub struct TestUserResponse {
pub pid: String,
pub user_id: i32,
pub user_email: String,
}
// Mock user struct for testing JWTWithUser extractor
#[derive(Debug, Clone)]
struct TestUser {
id: i32,
email: String,
}
#[async_trait::async_trait]
impl Authenticable for TestUser {
async fn find_by_claims_key(
_db: &sea_orm::DatabaseConnection,
pid: &str,
) -> Result<Self, ModelError> {
// Simple mock: return user if pid matches, otherwise not found
if pid == "test_pid_123" {
Ok(Self {
id: 1,
email: "test@example.com".to_string(),
})
} else {
Err(ModelError::EntityNotFound)
}
}
async fn find_by_api_key(
_db: &sea_orm::DatabaseConnection,
api_key: &str,
) -> Result<Self, ModelError> {
// Simple mock: return user if api_key matches, otherwise not found
if api_key == "test_api_key_123" {
Ok(Self {
id: 1,
email: "test@example.com".to_string(),
})
} else {
Err(ModelError::EntityNotFound)
}
}
}
// Test handler for JWTWithUser extractor
async fn jwt_with_user_handler(auth: auth::JWTWithUser<TestUser>) -> Result<Response> {
format::json(TestUserResponse {
pid: auth.claims.pid,
user_id: auth.user.id,
user_email: auth.user.email,
})
}
// Test JWTWithUser extractor with valid token
#[tokio::test]
async fn can_extract_jwt_with_user_valid_token() {
let mut ctx = tests_cfg::app::get_app_context().await;
// Configure JWT auth
let secret = "PqRwLF2rhHe8J22oBeHy".to_string();
ctx.config.auth = Some(loco_rs::config::Auth {
jwt: Some(loco_rs::config::JWT {
location: None,
secret: secret.clone(),
expiration: 3600,
}),
});
// Create a valid JWT token with known PID
let jwt = loco_rs::auth::jwt::JWT::new(&secret);
let token = jwt
.generate_token(3600, "test_pid_123".to_string(), serde_json::Map::new())
.expect("Failed to generate token");
let port = get_available_port().await;
let handle =
infra_cfg::server::start_with_route(ctx, "/", get(jwt_with_user_handler), Some(port)).await;
let client = reqwest::Client::new();
let res = client
.get(get_base_url_port(port))
.header("Authorization", format!("Bearer {token}"))
.send()
.await
.expect("Valid response");
assert_eq!(res.status(), 200);
let body: TestUserResponse = res.json().await.expect("Valid JSON response");
assert_eq!(body.pid, "test_pid_123");
assert_eq!(body.user_id, 1);
assert_eq!(body.user_email, "test@example.com");
handle.abort();
}
// Test JWTWithUser extractor with invalid token
#[tokio::test]
async fn can_handle_jwt_with_user_invalid_token() {
let mut ctx = tests_cfg::app::get_app_context().await;
// Configure JWT auth
let secret = "PqRwLF2rhHe8J22oBeHy".to_string();
ctx.config.auth = Some(loco_rs::config::Auth {
jwt: Some(loco_rs::config::JWT {
location: None,
secret: secret.clone(),
expiration: 3600,
}),
});
let port = get_available_port().await;
let handle =
infra_cfg::server::start_with_route(ctx, "/", get(jwt_with_user_handler), Some(port)).await;
let client = reqwest::Client::new();
let res = client
.get(get_base_url_port(port))
.header("Authorization", "Bearer invalid_token")
.send()
.await
.expect("Valid response");
assert_eq!(res.status(), 401);
handle.abort();
}
// Test JWTWithUser extractor with non-existent user
#[tokio::test]
async fn can_handle_jwt_with_user_nonexistent_user() {
let mut ctx = tests_cfg::app::get_app_context().await;
// Configure JWT auth
let secret = "PqRwLF2rhHe8J22oBeHy".to_string();
ctx.config.auth = Some(loco_rs::config::Auth {
jwt: Some(loco_rs::config::JWT {
location: None,
secret: secret.clone(),
expiration: 3600,
}),
});
// Create a valid JWT token with unknown PID
let jwt = loco_rs::auth::jwt::JWT::new(&secret);
let token = jwt
.generate_token(3600, "unknown_pid".to_string(), serde_json::Map::new())
.expect("Failed to generate token");
let port = get_available_port().await;
let handle =
infra_cfg::server::start_with_route(ctx, "/", get(jwt_with_user_handler), Some(port)).await;
let client = reqwest::Client::new();
let res = client
.get(get_base_url_port(port))
.header("Authorization", format!("Bearer {token}"))
.send()
.await
.expect("Valid response");
assert_eq!(res.status(), 401);
handle.abort();
}
// Test JWTWithUser extractor with missing token
#[tokio::test]
async fn can_handle_jwt_with_user_missing_token() {
let mut ctx = tests_cfg::app::get_app_context().await;
// Configure JWT auth
let secret = "PqRwLF2rhHe8J22oBeHy".to_string();
ctx.config.auth = Some(loco_rs::config::Auth {
jwt: Some(loco_rs::config::JWT {
location: None,
secret: secret.clone(),
expiration: 3600,
}),
});
let port = get_available_port().await;
let handle =
infra_cfg::server::start_with_route(ctx, "/", get(jwt_with_user_handler), Some(port)).await;
let client = reqwest::Client::new();
let res = client
.get(get_base_url_port(port))
.send()
.await
.expect("Valid response");
assert_eq!(res.status(), 401);
handle.abort();
}
@@ -0,0 +1,7 @@
mod jwt;
#[cfg(feature = "with-db")]
mod jwt_with_user;
#[cfg(feature = "with-db")]
mod api_token;
@@ -0,0 +1,4 @@
mod auth;
mod shared_store;
mod validate;
mod view_engine;
@@ -0,0 +1,87 @@
use axum::extract::State;
use loco_rs::{controller::format, prelude::*, tests_cfg};
use rstest::rstest;
use serde::{Deserialize, Serialize};
use crate::infra_cfg;
#[derive(Clone, Debug, Serialize, Deserialize, PartialEq)]
struct MySharedData {
message: String,
}
struct MySharedDataWithoutClone {
message: String,
}
#[rstest]
#[case(true)]
#[case(false)]
#[tokio::test]
async fn test_shared_store_extractor(#[case] exists: bool) {
async fn action(
State(_ctx): State<AppContext>,
SharedStore(shared_data): SharedStore<MySharedData>,
) -> Result<Response> {
format::json(&shared_data)
}
let ctx: AppContext = tests_cfg::app::get_app_context().await;
let test_data = MySharedData {
message: "Hello from SharedStore!".to_string(),
};
if exists {
ctx.shared_store.insert(test_data.clone());
}
let port = get_available_port().await;
let handle = infra_cfg::server::start_with_route(ctx, "/", get(action), Some(port)).await;
let res = reqwest::get(get_base_url_port(port))
.await
.expect("Failed to make request");
if exists {
assert_eq!(res.status(), axum::http::StatusCode::OK);
let body: MySharedData = res.json().await.expect("Failed to parse response body");
assert_eq!(body, test_data);
} else {
assert_eq!(res.status(), axum::http::StatusCode::INTERNAL_SERVER_ERROR);
}
handle.abort();
}
#[tokio::test]
async fn test_shared_store_without_clone() {
async fn action(State(ctx): State<AppContext>) -> Result<Response> {
let shared_data_ref = ctx
.shared_store
.get_ref::<MySharedDataWithoutClone>()
.ok_or_else(|| Error::InternalServerError)?;
format::text(&shared_data_ref.message)
}
let ctx: AppContext = tests_cfg::app::get_app_context().await;
let test_data = MySharedDataWithoutClone {
message: "Hello from SharedStore!".to_string(),
};
ctx.shared_store.insert(test_data);
let port = get_available_port().await;
let handle = infra_cfg::server::start_with_route(ctx, "/", get(action), Some(port)).await;
let res = reqwest::get(get_base_url_port(port))
.await
.expect("Failed to make request");
assert_eq!(res.status(), axum::http::StatusCode::OK);
let body = res.text().await.expect("Failed to parse response body");
assert_eq!(body, "Hello from SharedStore!");
handle.abort();
}
@@ -0,0 +1,90 @@
use loco_rs::{prelude::*, tests_cfg};
use serde::{Deserialize, Serialize};
use validator::Validate;
use crate::infra_cfg;
#[derive(Debug, Deserialize, Serialize, Validate)]
pub struct Data {
#[validate(length(min = 5, message = "message_str"))]
pub name: String,
#[validate(email)]
pub email: String,
}
async fn validation_with_response(
JsonValidateWithMessage(_params): JsonValidateWithMessage<Data>,
) -> Result<Response> {
format::json(())
}
async fn simple_validation(JsonValidate(_params): JsonValidate<Data>) -> Result<Response> {
format::json(())
}
#[tokio::test]
async fn can_validation_with_response() {
let ctx = tests_cfg::app::get_app_context().await;
let port = get_available_port().await;
let handle =
infra_cfg::server::start_with_route(ctx, "/", post(validation_with_response), Some(port))
.await;
let client = reqwest::Client::new();
let res = client
.post(get_base_url_port(port))
.json(&serde_json::json!({"name": "test", "email": "invalid"}))
.send()
.await
.expect("Valid response");
assert_eq!(res.status(), 400);
let res_text = res.text().await.expect("response text");
let res_json: serde_json::Value = serde_json::from_str(&res_text).expect("Valid JSON response");
let expected_json = serde_json::json!(
{
"errors":{
"email":[{"code":"email","message":null,"params":{"value":"invalid"}}],
"name":[{"code":"length","message":"message_str","params":{"min":5,"value":"test"}}]
}
});
assert_eq!(res_json, expected_json);
handle.abort();
}
#[tokio::test]
async fn can_validation_without_response() {
let ctx = tests_cfg::app::get_app_context().await;
let port = get_available_port().await;
let handle =
infra_cfg::server::start_with_route(ctx, "/", post(simple_validation), Some(port)).await;
let client = reqwest::Client::new();
let res = client
.post(get_base_url_port(port))
.json(&serde_json::json!({"name": "test", "email": "invalid"}))
.send()
.await
.expect("Valid response");
assert_eq!(res.status(), 400);
let res_text = res.text().await.expect("response text");
let res_json: serde_json::Value = serde_json::from_str(&res_text).expect("Valid JSON response");
let expected_json = serde_json::json!(
{
"error": "Bad Request"
}
);
assert_eq!(res_json, expected_json);
handle.abort();
}
@@ -0,0 +1,26 @@
use loco_rs::{prelude::*, tests_cfg};
use crate::infra_cfg;
async fn action(ViewEngine(_engine): ViewEngine<()>) -> Result<Response> {
format::json(())
}
/// When the `ViewEngine` layer (`Extension<ViewEngine<E>>`) was never
/// installed, the extractor must reject gracefully with an error response
/// instead of panicking.
#[tokio::test]
async fn missing_layer_rejects_gracefully() {
let ctx = tests_cfg::app::get_app_context().await;
let port = get_available_port().await;
let handle = infra_cfg::server::start_with_route(ctx, "/", get(action), Some(port)).await;
let res = reqwest::get(get_base_url_port(port))
.await
.expect("valid response");
assert_eq!(res.status(), axum::http::StatusCode::INTERNAL_SERVER_ERROR);
handle.abort();
}
@@ -0,0 +1,95 @@
use axum::extract::FromRef;
use loco_rs::{
app::{AppContext, SharedStore},
cache,
prelude::*,
tests_cfg,
};
use std::sync::Arc;
use crate::infra_cfg;
#[cfg(feature = "with-db")]
use sea_orm::DatabaseConnection;
/// Tests that DatabaseConnection can be extracted from AppContext via FromRef
#[cfg(feature = "with-db")]
#[tokio::test]
async fn can_extract_db_connection_from_app_context() {
let ctx = tests_cfg::app::get_app_context().await;
#[allow(clippy::items_after_statements)]
async fn action(State(ctx): State<AppContext>) -> Result<Response> {
// Use FromRef to extract DatabaseConnection from AppContext
let _db: DatabaseConnection = DatabaseConnection::from_ref(&ctx);
format::json(serde_json::json!({"extracted": "db"}))
}
let port = get_available_port().await;
let handle = infra_cfg::server::start_with_route(ctx, "/", get(action), Some(port)).await;
let res = reqwest::get(get_base_url_port(port))
.await
.expect("Valid response");
assert_eq!(res.status(), 200);
let body: serde_json::Value = res.json().await.expect("JSON response");
assert_eq!(body["extracted"], "db");
handle.abort();
}
/// Tests that Arc<Cache> can be extracted from AppContext via FromRef
#[tokio::test]
async fn can_extract_cache_from_app_context() {
let ctx = tests_cfg::app::get_app_context().await;
#[allow(clippy::items_after_statements)]
async fn action(State(ctx): State<AppContext>) -> Result<Response> {
// Use FromRef to extract Arc<Cache> from AppContext
let _cache: Arc<cache::Cache> = Arc::from_ref(&ctx);
format::json(serde_json::json!({"extracted": "cache"}))
}
let port = get_available_port().await;
let handle = infra_cfg::server::start_with_route(ctx, "/", get(action), Some(port)).await;
let res = reqwest::get(get_base_url_port(port))
.await
.expect("Valid response");
assert_eq!(res.status(), 200);
let body: serde_json::Value = res.json().await.expect("JSON response");
assert_eq!(body["extracted"], "cache");
handle.abort();
}
/// Tests that Arc<SharedStore> can be extracted from AppContext via FromRef
#[tokio::test]
async fn can_extract_shared_store_from_app_context() {
let ctx = tests_cfg::app::get_app_context().await;
#[allow(clippy::items_after_statements)]
async fn action(State(ctx): State<AppContext>) -> Result<Response> {
// Use FromRef to extract Arc<SharedStore> from AppContext
let _store: Arc<SharedStore> = Arc::from_ref(&ctx);
format::json(serde_json::json!({"extracted": "shared_store"}))
}
let port = get_available_port().await;
let handle = infra_cfg::server::start_with_route(ctx, "/", get(action), Some(port)).await;
let res = reqwest::get(get_base_url_port(port))
.await
.expect("Valid response");
assert_eq!(res.status(), 200);
let body: serde_json::Value = res.json().await.expect("JSON response");
assert_eq!(body["extracted"], "shared_store");
handle.abort();
}
@@ -0,0 +1,206 @@
use loco_rs::{controller, prelude::*, tests_cfg};
use serde::{Deserialize, Serialize};
use crate::infra_cfg;
#[tokio::test]
async fn not_found() {
let ctx = tests_cfg::app::get_app_context().await;
#[allow(clippy::items_after_statements)]
async fn action() -> Result<Response> {
controller::not_found()
}
let port = get_available_port().await;
let handle = infra_cfg::server::start_with_route(ctx, "/", get(action), Some(port)).await;
let res = reqwest::get(get_base_url_port(port))
.await
.expect("Valid response");
assert_eq!(res.status(), 404);
let res_text = res.text().await.expect("response text");
let res_json: serde_json::Value = serde_json::from_str(&res_text).expect("Valid JSON response");
let expected_json = serde_json::json!({
"error": "not_found",
"description": "Resource was not found"
});
assert_eq!(res_json, expected_json);
handle.abort();
}
#[tokio::test]
async fn internal_server_error() {
let ctx = tests_cfg::app::get_app_context().await;
#[allow(clippy::items_after_statements)]
async fn action() -> Result<Response> {
Err(Error::InternalServerError)
}
let port = get_available_port().await;
let handle = infra_cfg::server::start_with_route(ctx, "/", get(action), Some(port)).await;
let res = reqwest::get(get_base_url_port(port))
.await
.expect("Valid response");
assert_eq!(res.status(), 500);
let res_text = res.text().await.expect("response text");
let res_json: serde_json::Value = serde_json::from_str(&res_text).expect("Valid JSON response");
let expected_json = serde_json::json!({
"error": "internal_server_error",
"description": "Internal Server Error",
});
assert_eq!(res_json, expected_json);
handle.abort();
}
#[tokio::test]
async fn unauthorized() {
let ctx = tests_cfg::app::get_app_context().await;
#[allow(clippy::items_after_statements)]
async fn action() -> Result<Response> {
controller::unauthorized("user not unauthorized")
}
let port = get_available_port().await;
let handle = infra_cfg::server::start_with_route(ctx, "/", get(action), Some(port)).await;
let res = reqwest::get(get_base_url_port(port))
.await
.expect("Valid response");
assert_eq!(res.status(), 401);
let res_text = res.text().await.expect("response text");
let res_json: serde_json::Value = serde_json::from_str(&res_text).expect("Valid JSON response");
let expected_json = serde_json::json!({
"error": "unauthorized",
"description": "You do not have permission to access this resource"
});
assert_eq!(res_json, expected_json);
handle.abort();
}
#[tokio::test]
async fn fallback() {
let ctx = tests_cfg::app::get_app_context().await;
#[allow(clippy::items_after_statements)]
async fn action() -> Result<Response> {
Err(Error::Message(String::new()))
}
let port = get_available_port().await;
let handle = infra_cfg::server::start_with_route(ctx, "/", get(action), Some(port)).await;
let res = reqwest::get(get_base_url_port(port))
.await
.expect("Valid response");
assert_eq!(res.status(), 500);
let res_text = res.text().await.expect("response text");
let res_json: serde_json::Value = serde_json::from_str(&res_text).expect("Valid JSON response");
let expected_json = serde_json::json!({
"error": "internal_server_error",
"description": "Internal Server Error",
});
assert_eq!(res_json, expected_json);
handle.abort();
}
#[tokio::test]
async fn custom_error() {
let ctx = tests_cfg::app::get_app_context().await;
#[allow(clippy::items_after_statements)]
async fn action() -> Result<Response> {
Err(Error::CustomError(
axum::http::StatusCode::PAYLOAD_TOO_LARGE,
controller::ErrorDetail {
error: Some("Payload Too Large".to_string()),
description: Some("413 Payload Too Large".to_string()),
errors: None,
},
))
}
let port = get_available_port().await;
let handle = infra_cfg::server::start_with_route(ctx, "/", get(action), Some(port)).await;
let res = reqwest::get(get_base_url_port(port))
.await
.expect("Valid response");
assert_eq!(res.status(), 413);
let res_text = res.text().await.expect("response text");
let res_json: serde_json::Value = serde_json::from_str(&res_text).expect("Valid JSON response");
let expected_json = serde_json::json!({
"error": "Payload Too Large",
"description": "413 Payload Too Large"
});
assert_eq!(res_json, expected_json);
handle.abort();
}
#[tokio::test]
async fn json_rejection() {
let ctx = tests_cfg::app::get_app_context().await;
#[allow(clippy::items_after_statements)]
#[derive(Debug, Deserialize, Serialize)]
pub struct Data {
pub email: String,
}
#[allow(clippy::items_after_statements)]
async fn action(Json(_params): Json<Data>) -> Result<Response> {
format::json(())
}
let port = get_available_port().await;
let handle = infra_cfg::server::start_with_route(ctx, "/", post(action), Some(port)).await;
let client = reqwest::Client::new();
let res = client
.post(get_base_url_port(port))
.json(&serde_json::json!({}))
.send()
.await
.expect("Valid response");
assert_eq!(res.status(), 422);
let res_text = res.text().await.expect("response text");
let res_json: serde_json::Value = serde_json::from_str(&res_text).expect("Valid JSON response");
let expected_json = serde_json::json!({
"error": "Bad Request",
});
assert_eq!(res_json, expected_json);
handle.abort();
}
@@ -0,0 +1,489 @@
use std::{collections::BTreeMap, path::PathBuf};
use axum::http::StatusCode;
use insta::assert_debug_snapshot;
use loco_rs::{controller::middleware, prelude::*, tests_cfg};
use rstest::rstest;
use crate::infra_cfg;
macro_rules! configure_insta {
($($expr:expr_2021),*) => {
let mut settings = insta::Settings::clone_current();
settings.set_prepend_module_to_snapshot(false);
settings.set_snapshot_suffix("middlewares");
let _guard = settings.bind_to_scope();
};
}
#[rstest]
#[case(true)]
#[case(false)]
#[tokio::test]
async fn panic(#[case] enable: bool) {
configure_insta!();
#[allow(clippy::items_after_statements)]
async fn action() -> Result<Response> {
panic!("panic!")
}
let mut ctx: AppContext = tests_cfg::app::get_app_context().await;
ctx.config.server.middlewares.catch_panic =
Some(middleware::catch_panic::CatchPanic { enable });
let port = get_available_port().await;
let handle = infra_cfg::server::start_with_route(ctx, "/", get(action), Some(port)).await;
let res = reqwest::get(get_base_url_port(port)).await;
if enable {
let res = res.expect("valid response");
assert_debug_snapshot!(
format!("panic"),
(res.status().to_string(), res.text().await)
);
} else {
assert!(res.is_err());
}
handle.abort();
}
#[rstest]
#[case(true)]
#[case(false)]
#[tokio::test]
async fn etag(#[case] enable: bool) {
async fn action() -> Result<Response> {
format::render().etag("loco-etag")?.text("content")
}
let mut ctx: AppContext = tests_cfg::app::get_app_context().await;
ctx.config.server.middlewares.etag = Some(middleware::etag::Etag { enable });
let port = get_available_port().await;
let handle = infra_cfg::server::start_with_route(ctx, "/", get(action), Some(port)).await;
let res = reqwest::Client::new()
.get(get_base_url_port(port))
.header("if-none-match", "loco-etag")
.send()
.await
.expect("response");
if enable {
assert_eq!(res.status(), StatusCode::NOT_MODIFIED);
} else {
assert_eq!(res.status(), StatusCode::OK);
}
handle.abort();
}
#[rstest]
// Enabled with the default `RightmostXForwardedFor` source: the rightmost
// (last) value of `X-Forwarded-For` is taken verbatim, with **no**
// trusted-proxy CIDR skipping (that behavior no longer exists) — so the
// last hop (`192.1.1.1` here) wins, even though the old middleware would
// have treated it as a trusted proxy and skipped it.
#[case(true, None, "remote: 192.1.1.1")]
// Configuring an alternative source picks a different header entirely.
#[case(
true,
Some(middleware::remote_ip::ClientIpSource::XRealIp),
"remote: 9.9.9.9"
)]
// Disabled: no `ClientIpSource` extension is inserted into the router at
// all, so the extractor resolves to `RemoteIP::None`.
#[case(false, None, "--")]
#[tokio::test]
async fn remote_ip(
#[case] enable: bool,
#[case] source: Option<middleware::remote_ip::ClientIpSource>,
#[case] expected: &str,
) {
#[allow(clippy::items_after_statements)]
async fn action(remote_ip: RemoteIP) -> Result<Response> {
format::text(&remote_ip.to_string())
}
let mut ctx: AppContext = tests_cfg::app::get_app_context().await;
let mut middleware_config = middleware::remote_ip::RemoteIpMiddleware {
enable,
..Default::default()
};
if let Some(source) = source {
middleware_config.source = source;
}
ctx.config.server.middlewares.remote_ip = Some(middleware_config);
let port = get_available_port().await;
let handle = infra_cfg::server::start_with_route(ctx, "/", get(action), Some(port)).await;
let res = reqwest::Client::new()
.get(get_base_url_port(port))
.header(
"x-forwarded-for",
reqwest::header::HeaderValue::from_static("51.50.51.50,192.1.1.1"),
)
.header(
"x-real-ip",
reqwest::header::HeaderValue::from_static("9.9.9.9"),
)
.send()
.await
.expect("response");
assert_eq!(res.text().await.expect("string"), expected.to_string());
handle.abort();
}
#[rstest]
#[case(true)]
#[case(false)]
#[tokio::test]
async fn timeout(#[case] enable: bool) {
#[allow(clippy::items_after_statements)]
async fn action() -> Result<Response> {
tokio::time::sleep(tokio::time::Duration::from_secs(3)).await;
format::render().text("loco")
}
let mut ctx: AppContext = tests_cfg::app::get_app_context().await;
ctx.config.server.middlewares.timeout_request =
Some(middleware::timeout::TimeOut { enable, timeout: 2 });
let port = get_available_port().await;
let handle = infra_cfg::server::start_with_route(ctx, "/", get(action), Some(port)).await;
let res = reqwest::get(get_base_url_port(port))
.await
.expect("response");
if enable {
assert_eq!(res.status(), StatusCode::REQUEST_TIMEOUT);
} else {
assert_eq!(res.status(), StatusCode::OK);
}
handle.abort();
}
#[rstest]
#[case(true, "default", None, None, None)]
#[case(true, "with_allow_headers", Some(vec!["token".to_string(), "user".to_string()]), None, None)]
#[case(true, "with_allow_methods", None, Some(vec!["post".to_string(), "get".to_string()]), None)]
#[case(true, "with_max_age", None, None, Some(20))]
#[case(false, "disabled", None, None, None)]
#[tokio::test]
async fn cors(
#[case] enable: bool,
#[case] test_name: &str,
#[case] allow_headers: Option<Vec<String>>,
#[case] allow_methods: Option<Vec<String>>,
#[case] max_age: Option<u64>,
) {
use loco_rs::controller::middleware::cors::Cors;
configure_insta!();
let mut ctx: AppContext = tests_cfg::app::get_app_context().await;
let mut middleware = Cors {
enable,
..Default::default()
};
if let Some(allow_headers) = allow_headers {
middleware.allow_headers = allow_headers;
}
if let Some(allow_methods) = allow_methods {
middleware.allow_methods = allow_methods;
}
middleware.max_age = max_age;
ctx.config.server.middlewares.cors = Some(middleware);
let port = get_available_port().await;
let handle = infra_cfg::server::start_from_ctx(ctx, Some(port)).await;
let res = reqwest::Client::new()
.request(reqwest::Method::OPTIONS, get_base_url_port(port))
.send()
.await
.expect("valid response");
assert_debug_snapshot!(
format!("cors_[{test_name}]"),
(
format!(
"access-control-allow-origin: {:?}",
res.headers().get("access-control-allow-origin")
),
format!("vary: {:?}", res.headers().get("vary")),
format!(
"access-control-allow-methods: {:?}",
res.headers().get("access-control-allow-methods")
),
format!(
"access-control-allow-headers: {:?}",
res.headers().get("access-control-allow-headers")
),
format!("allow: {:?}", res.headers().get("allow")),
)
);
handle.abort();
}
#[rstest]
#[case(middleware::limit_payload::DefaultBodyLimitKind::Limit(0x1B))]
#[case(middleware::limit_payload::DefaultBodyLimitKind::Disable)]
#[tokio::test]
async fn limit_payload(#[case] limit: middleware::limit_payload::DefaultBodyLimitKind) {
configure_insta!();
let mut ctx: AppContext = tests_cfg::app::get_app_context().await;
ctx.config.server.middlewares.limit_payload =
Some(middleware::limit_payload::LimitPayload { body_limit: limit });
let port = get_available_port().await;
let handle = infra_cfg::server::start_from_ctx(ctx, Some(port)).await;
let res = reqwest::Client::new()
.request(reqwest::Method::POST, get_base_url_port(port))
.body("send body".repeat(100))
.send()
.await
.expect("valid response");
match limit {
middleware::limit_payload::DefaultBodyLimitKind::Disable => {
assert_eq!(res.status(), StatusCode::OK);
}
middleware::limit_payload::DefaultBodyLimitKind::Limit(_) => {
assert_eq!(res.status(), StatusCode::PAYLOAD_TOO_LARGE);
}
}
handle.abort();
}
#[cfg(not(feature = "embedded_assets"))]
#[tokio::test]
async fn static_assets() {
configure_insta!();
let base_static_assets_path = PathBuf::from("assets").join("static");
let static_asset_path = tree_fs::TreeBuilder::default()
.drop(true)
.add(
base_static_assets_path.join("404.html"),
"<h1>404 not found</h1>",
)
.add(
base_static_assets_path.join("static.html"),
"<h1>static content</h1>",
)
.create()
.expect("create static tree file");
let mut ctx: AppContext = tests_cfg::app::get_app_context().await;
let base_static_path = static_asset_path.root.join(base_static_assets_path);
ctx.config.server.middlewares.static_assets = Some(middleware::static_assets::StaticAssets {
enable: true,
must_exist: true,
folder: middleware::static_assets::FolderConfig {
uri: "/static".to_string(),
path: base_static_path.clone(),
},
fallback: base_static_path.join("404.html"),
precompressed: false,
cache_control: None,
});
let port = get_available_port().await;
let handle = infra_cfg::server::start_from_ctx(ctx, Some(port)).await;
let get_static_html = reqwest::get(format!("{}static/static.html", get_base_url_port(port)))
.await
.expect("valid response");
assert_eq!(
get_static_html.text().await.expect("text response"),
"<h1>static content</h1>".to_string()
);
let get_fallback = reqwest::get(format!("{}static/logo.png", get_base_url_port(port)))
.await
.expect("valid response");
assert_eq!(
get_fallback.text().await.expect("text response"),
"<h1>404 not found</h1>".to_string()
);
handle.abort();
}
#[rstest]
#[case(None, None)]
#[case(Some("empty".to_string()), None)]
#[case(Some("github".to_string()), Some(BTreeMap::from([(
"Content-Security-Policy".to_string(),
"default-src 'self' https".to_string(),
)])))]
#[tokio::test]
async fn secure_headers(
#[case] preset: Option<String>,
#[case] overrides: Option<BTreeMap<String, String>>,
) {
configure_insta!();
let mut ctx: AppContext = tests_cfg::app::get_app_context().await;
ctx.config.server.middlewares.secure_headers = Some(
loco_rs::controller::middleware::secure_headers::SecureHeader {
enable: true,
preset: preset.clone().unwrap_or_else(|| "github".to_string()),
overrides: overrides.clone(),
},
);
let port = get_available_port().await;
let handle = infra_cfg::server::start_from_ctx(ctx, Some(port)).await;
let res = reqwest::Client::new()
.request(reqwest::Method::POST, get_base_url_port(port))
.send()
.await
.expect("response");
let policy = res.headers().get("content-security-policy");
let overrides_str = overrides.map_or("none".to_string(), |k| {
k.keys()
.map(std::string::ToString::to_string)
.collect::<Vec<_>>()
.join(",")
});
assert_debug_snapshot!(
format!(
"secure_headers_[{}]_overrides[{}]",
preset.unwrap_or_else(|| "none".to_string()),
overrides_str
),
policy
);
handle.abort();
}
#[rstest]
#[case(None, false, None)]
#[case(Some(StatusCode::BAD_REQUEST), false, None)]
#[case(None, true, None)]
#[case(None, false, Some("text fallback response".to_string()))]
#[tokio::test]
async fn fallback(
#[case] code: Option<StatusCode>,
#[case] file: bool,
#[case] not_found: Option<String>,
) {
let mut ctx: AppContext = tests_cfg::app::get_app_context().await;
let maybe_file = if file {
Some(
tree_fs::TreeBuilder::default()
.drop(true)
.add(
PathBuf::from("static_content.html"),
"<h1>fallback response</h1>",
)
.create()
.unwrap(),
)
} else {
None
};
let mut fallback_config = middleware::fallback::Fallback {
enable: true,
file: maybe_file.as_ref().map(|tree_fs| {
tree_fs
.root
.join("static_content.html")
.display()
.to_string()
}),
not_found: not_found.clone(),
..Default::default()
};
if let Some(code) = code {
fallback_config.code = code;
};
ctx.config.server.middlewares.fallback = Some(fallback_config);
let port = get_available_port().await;
let handle = infra_cfg::server::start_from_ctx(ctx, Some(port)).await;
let res = reqwest::get(format!("{}not-found", get_base_url_port(port)))
.await
.expect("valid response");
if let Some(code) = code {
assert_eq!(res.status(), code);
} else if maybe_file.is_some() {
// the file fallback is served via `ServeFile`, which reports its own
// status (200 OK for a found file) regardless of the configured
// `code`.
assert_eq!(res.status(), StatusCode::OK);
} else {
assert_eq!(res.status(), StatusCode::NOT_FOUND);
}
let response_text = res.text().await.expect("response text");
if maybe_file.is_some() {
assert_eq!(response_text, "<h1>fallback response</h1>".to_string());
}
if let Some(not_found_text) = not_found {
assert_eq!(response_text, not_found_text);
}
handle.abort();
}
#[rstest]
#[case(None)]
#[case(Some("custom".to_string()))]
#[tokio::test]
async fn powered_by_header(#[case] ident: Option<String>) {
configure_insta!();
let mut ctx: AppContext = tests_cfg::app::get_app_context().await;
ctx.config.server.ident.clone_from(&ident);
let port = get_available_port().await;
let handle = infra_cfg::server::start_from_ctx(ctx, Some(port)).await;
let res = reqwest::get(get_base_url_port(port))
.await
.expect("valid response");
let header_value = res.headers().get("x-powered-by").expect("exists header");
if let Some(ident_str) = ident {
assert_eq!(header_value.to_str().expect("value"), ident_str);
} else {
assert_eq!(header_value.to_str().expect("value"), "loco.rs");
}
handle.abort();
}
+4
View File
@@ -0,0 +1,4 @@
mod extractor;
mod from_ref;
mod into_response;
mod middlewares;
@@ -0,0 +1,11 @@
---
source: tests/controller/middlewares.rs
expression: "(format!(\"access-control-allow-origin: {:?}\",\n res.headers().get(\"access-control-allow-origin\")),\n format!(\"vary: {:?}\", res.headers().get(\"vary\")),\n format!(\"access-control-allow-methods: {:?}\",\n res.headers().get(\"access-control-allow-methods\")),\n format!(\"access-control-allow-headers: {:?}\",\n res.headers().get(\"access-control-allow-headers\")),\n format!(\"allow: {:?}\", res.headers().get(\"allow\")))"
---
(
"access-control-allow-origin: Some(\"*\")",
"vary: Some(\"origin, access-control-request-method, access-control-request-headers\")",
"access-control-allow-methods: Some(\"*\")",
"access-control-allow-headers: Some(\"*\")",
"allow: Some(\"GET,HEAD,POST\")",
)
@@ -0,0 +1,11 @@
---
source: tests/controller/middlewares.rs
expression: "(format!(\"access-control-allow-origin: {:?}\",\n res.headers().get(\"access-control-allow-origin\")),\n format!(\"vary: {:?}\", res.headers().get(\"vary\")),\n format!(\"access-control-allow-methods: {:?}\",\n res.headers().get(\"access-control-allow-methods\")),\n format!(\"access-control-allow-headers: {:?}\",\n res.headers().get(\"access-control-allow-headers\")),\n format!(\"allow: {:?}\", res.headers().get(\"allow\")))"
---
(
"access-control-allow-origin: None",
"vary: None",
"access-control-allow-methods: None",
"access-control-allow-headers: None",
"allow: Some(\"GET,HEAD,POST\")",
)
@@ -0,0 +1,11 @@
---
source: tests/controller/middlewares.rs
expression: "(format!(\"access-control-allow-origin: {:?}\",\n res.headers().get(\"access-control-allow-origin\")),\n format!(\"vary: {:?}\", res.headers().get(\"vary\")),\n format!(\"access-control-allow-methods: {:?}\",\n res.headers().get(\"access-control-allow-methods\")),\n format!(\"access-control-allow-headers: {:?}\",\n res.headers().get(\"access-control-allow-headers\")),\n format!(\"allow: {:?}\", res.headers().get(\"allow\")))"
---
(
"access-control-allow-origin: Some(\"*\")",
"vary: Some(\"origin, access-control-request-method, access-control-request-headers\")",
"access-control-allow-methods: Some(\"*\")",
"access-control-allow-headers: Some(\"token,user\")",
"allow: Some(\"GET,HEAD,POST\")",
)
@@ -0,0 +1,11 @@
---
source: tests/controller/middlewares.rs
expression: "(format!(\"access-control-allow-origin: {:?}\",\n res.headers().get(\"access-control-allow-origin\")),\n format!(\"vary: {:?}\", res.headers().get(\"vary\")),\n format!(\"access-control-allow-methods: {:?}\",\n res.headers().get(\"access-control-allow-methods\")),\n format!(\"access-control-allow-headers: {:?}\",\n res.headers().get(\"access-control-allow-headers\")),\n format!(\"allow: {:?}\", res.headers().get(\"allow\")))"
---
(
"access-control-allow-origin: Some(\"*\")",
"vary: Some(\"origin, access-control-request-method, access-control-request-headers\")",
"access-control-allow-methods: Some(\"post,get\")",
"access-control-allow-headers: Some(\"*\")",
"allow: Some(\"GET,HEAD,POST\")",
)
@@ -0,0 +1,11 @@
---
source: tests/controller/middlewares.rs
expression: "(format!(\"access-control-allow-origin: {:?}\",\n res.headers().get(\"access-control-allow-origin\")),\n format!(\"vary: {:?}\", res.headers().get(\"vary\")),\n format!(\"access-control-allow-methods: {:?}\",\n res.headers().get(\"access-control-allow-methods\")),\n format!(\"access-control-allow-headers: {:?}\",\n res.headers().get(\"access-control-allow-headers\")),\n format!(\"allow: {:?}\", res.headers().get(\"allow\")))"
---
(
"access-control-allow-origin: Some(\"*\")",
"vary: Some(\"origin, access-control-request-method, access-control-request-headers\")",
"access-control-allow-methods: Some(\"*\")",
"access-control-allow-headers: Some(\"*\")",
"allow: Some(\"GET,HEAD,POST\")",
)
@@ -0,0 +1,10 @@
---
source: tests/controller/middlewares.rs
expression: "(res.status().to_string(), res.text().await)"
---
(
"500 Internal Server Error",
Ok(
"{\"error\":\"internal_server_error\",\"description\":\"Internal Server Error\"}",
),
)
@@ -0,0 +1,5 @@
---
source: tests/controller/middlewares.rs
expression: policy
---
None
@@ -0,0 +1,7 @@
---
source: tests/controller/middlewares.rs
expression: policy
---
Some(
"default-src 'self' https",
)
@@ -0,0 +1,7 @@
---
source: tests/controller/middlewares.rs
expression: policy
---
Some(
"default-src 'self' https:; font-src 'self' https: data:; img-src 'self' https: data:; object-src 'none'; script-src https:; style-src 'self' https: 'unsafe-inline'",
)
@@ -0,0 +1,7 @@
<!DOCTYPE html>
<html>
<head><title>{% block title %}Email{% endblock %}</title></head>
<body>
{% block body %}{% endblock %}
</body>
</html>
@@ -0,0 +1,6 @@
{% extends "base.t" %}
{% block title %}Welcome Email{% endblock %}
{% block body %}
<h1>Hello {{ name }}!</h1>
<p>Your verification token is: <strong>{{ verifyToken }}</strong></p>
{% endblock %}
@@ -0,0 +1 @@
Welcome {{ name }}!
@@ -0,0 +1,5 @@
Hello {{ name }}!
Your verification token is: {{ verifyToken }}
Thank you for using our service.
@@ -0,0 +1,4 @@
{% extends "non_existent_base.t" %}
{% block content %}
<h1>This should fail</h1>
{% endblock %}
@@ -0,0 +1 @@
Subject template
@@ -0,0 +1 @@
Text template
@@ -0,0 +1,7 @@
{% extends "base.t" %}
{% block title %}Reset Password{% endblock %}
{% block body %}
<h1>Reset Your Password</h1>
<p>Click the link below to reset your password:</p>
<p><a href="{{ resetUrl }}">Reset Password</a></p>
{% endblock %}
@@ -0,0 +1 @@
Reset Your Password
@@ -0,0 +1,4 @@
Reset Your Password
Click the link below to reset your password:
{{ resetUrl }}
@@ -0,0 +1,16 @@
<!DOCTYPE html>
<html>
<head>
<title>{% block title %}Email{% endblock %}</title>
<style>
body { font-family: Arial, sans-serif; }
.footer { color: #666; font-size: 12px; }
</style>
</head>
<body>
{% block body %}{% endblock %}
<div class="footer">
{% block footer %}© 2024 My Company{% endblock %}
</div>
</body>
</html>
@@ -0,0 +1 @@
{% block subject %}Default Subject{% endblock %}
@@ -0,0 +1 @@
{% block text %}Default text content{% endblock %}
@@ -0,0 +1,10 @@
;<html>
<body>
This is a test content
<a href="http://localhost:/verify/{{ verifyToken }}">
Some test
</a>
</body>
</html>
@@ -0,0 +1 @@
Test {{ name }}
@@ -0,0 +1,3 @@
Welcome to test: {{ name }},
http://localhost/verify/<%= verifyToken %>
@@ -0,0 +1,6 @@
{% extends "base.t" %}
{% block title %}Welcome!{% endblock %}
{% block body %}
<h1>Welcome {{ name }}!</h1>
<p>Thank you for joining us.</p>
{% endblock %}
@@ -0,0 +1 @@
Welcome {{ name }}!
@@ -0,0 +1,3 @@
Welcome {{ name }}!
Thank you for joining us.
@@ -0,0 +1,153 @@
- id: "01JDM0X8EVAM823JZBGKYNBA99"
name: "UserAccountActivation"
task_data:
user_id: 133
email: "user11@example.com"
activation_token: "abcdef123456"
status: "queued"
run_at: "2024-11-28T08:19:08Z"
created_at: "2024-11-28T08:03:25Z"
updated_at: "2024-11-28T08:03:25Z"
- id: "01JDM0X8EVAM823JZBGKYNBA98"
name: "PasswordChangeNotification"
task_data:
user_id: 134
email: "user12@example.com"
change_time: "2024-11-27T12:30:00Z"
status: "completed"
run_at: "2024-11-28T08:04:25Z"
created_at: "2024-11-28T08:03:25Z"
updated_at: "2024-11-28T08:03:25Z"
- id: "01JDM0X8EVAM823JZBGKYNBA97"
name: "SendInvoice"
task_data:
user_id: 135
email: "user13@example.com"
invoice_id: "INV-2024-01"
status: "processing"
run_at: "2024-11-28T08:04:25Z"
created_at: "2024-11-28T08:03:25Z"
updated_at: "2024-11-28T08:03:25Z"
- id: "01JDM0X8EVAM823JZBGKYNBA96"
name: "UserDeactivation"
task_data:
user_id: 136
email: "user14@example.com"
deactivation_reason: "user requested"
status: "failed"
run_at: "2024-11-28T08:04:25Z"
created_at: "2024-11-28T08:03:25Z"
updated_at: "2024-11-28T08:03:25Z"
- id: "01JDM0X8EVAM823JZBGKYNBA95"
name: "SubscriptionReminder"
task_data:
user_id: 137
email: "user15@example.com"
renewal_date: "2024-12-01"
status: "queued"
run_at: "2024-11-28T08:04:25Z"
created_at: "2024-11-28T08:03:25Z"
updated_at: "2024-11-28T08:03:25Z"
- id: "01JDM0X8EVAM823JZBGKYNBA94"
name: "DataBackup"
task_data:
backup_id: "backup-12345"
user_id: 138
email: "user16@example.com"
status: "cancelled"
run_at: "2024-11-28T08:04:25Z"
created_at: "2024-11-28T08:03:25Z"
updated_at: "2024-11-28T08:03:25Z"
- id: "01JDM0X8EVAM823JZBGKYNBA93"
name: "SecurityAlert"
task_data:
user_id: 139
email: "user17@example.com"
alert_type: "login attempt from new device"
status: "queued"
run_at: "2024-11-28T08:04:25Z"
created_at: "2024-11-28T08:03:25Z"
updated_at: "2024-11-28T08:03:25Z"
- id: "01JDM0X8EVAM823JZBGKYNBA92"
name: "WeeklyReportEmail"
task_data:
user_id: 140
email: "user18@example.com"
report_period: "2024-11-20 to 2024-11-27"
status: "processing"
run_at: "2024-11-28T08:04:25Z"
created_at: "2024-11-28T08:03:25Z"
updated_at: "2024-11-28T08:03:25Z"
- id: "01JDM0X8EVAM823JZBGKYNBA91"
name: "AccountDeletion"
task_data:
user_id: 142
email: "user20@example.com"
deletion_request_time: "2024-11-27T14:00:00Z"
status: "queued"
run_at: "2024-11-28T08:04:25Z"
created_at: "2024-11-28T08:03:25Z"
updated_at: "2024-11-28T08:03:25Z"
- id: "01JDM0X8EVAM823JZBGKYNBA90"
name: "UserAccountActivation"
task_data:
user_id: 143
email: "user21@example.com"
activation_token: "xyz987654"
status: "completed"
run_at: "2024-11-28T08:04:25Z"
created_at: "2024-11-28T08:03:25Z"
updated_at: "2024-11-28T08:03:25Z"
- id: "01JDM0X8EVAM823JZBGKYNBA89"
name: "PasswordChangeNotification"
task_data:
user_id: 144
email: "user22@example.com"
change_time: "2024-11-27T15:00:00Z"
status: "completed"
run_at: "2024-11-28T08:04:25Z"
created_at: "2024-11-28T08:03:25Z"
updated_at: "2024-11-28T08:03:25Z"
- id: "01JDM0X8EVAM823JZBGKYNBA88"
name: "SendInvoice"
task_data:
user_id: 145
email: "user23@example.com"
invoice_id: "INV-2024-02"
status: "processing"
run_at: "2024-11-28T08:04:25Z"
created_at: "2024-11-28T08:03:25Z"
updated_at: "2024-11-28T08:03:25Z"
- id: "01JDM0X8EVAM823JZBGKYNBA87"
name: "UserDeactivation"
task_data:
user_id: 146
email: "user24@example.com"
deactivation_reason: "account inactive"
status: "failed"
run_at: "2024-11-28T08:04:25Z"
created_at: "2024-11-28T08:03:25Z"
updated_at: "2024-11-28T08:03:25Z"
- id: "01JDM0X8EVAM823JZBGKYNBA86"
name: "SubscriptionReminder"
task_data:
user_id: 147
email: "user25@example.com"
renewal_date: "2024-12-05"
status: "queued"
run_at: "2024-11-28T08:04:25Z"
created_at: "2024-11-28T08:03:25Z"
updated_at: "2024-11-28T08:03:25Z"
+1
View File
@@ -0,0 +1 @@
pub mod server;
+106
View File
@@ -0,0 +1,106 @@
//! # Server Infrastructure Utilities for Loco Framework Testing
//!
//! This module provides utility functions to test a server using the Loco
//! framework. It includes helper functions to start the server from different
//! configurations, such as from boot parameters, application context, or a
//! custom route. These utilities are designed for test environments and use
//! hardcoded ports and bindings.
use loco_rs::{boot, controller::AppRoutes, prelude::*, tests_cfg::db::AppHook};
/// A simple asynchronous handler for GET requests.
async fn get_action() -> Result<Response> {
format::render().text("text response")
}
/// A simple asynchronous handler for POST requests.
async fn post_action(_body: axum::body::Bytes) -> Result<Response> {
format::render().text("text response")
}
/// Starts the server using the provided Loco [`boot::BootResult`] result.
/// It uses hardcoded server parameters such as the port and binding address.
///
/// After spawning the server task, this polls the bound address until it
/// accepts a TCP connection (or a bounded timeout elapses), so callers can
/// issue requests the instant the listener is actually up — no fixed sleep.
pub async fn start_from_boot(
boot_result: boot::BootResult,
port: Option<i32>,
) -> tokio::task::JoinHandle<()> {
let port = port.unwrap_or(TEST_PORT_SERVER);
let handle = tokio::spawn(async move {
boot::start::<AppHook>(
boot_result,
boot::ServeParams {
port,
binding: TEST_BINDING_SERVER.to_string(),
},
false,
)
.await
.expect("start the server");
});
wait_until_ready(TEST_BINDING_SERVER, port).await;
handle
}
/// Polls `binding:port` until it accepts a TCP connection so a test can proceed
/// the moment the server is listening, replacing a fixed-duration sleep. Gives
/// up after a bounded number of attempts (~5s) rather than hanging a broken
/// boot forever — a failure then surfaces as the test's own request erroring,
/// which is the correct signal.
async fn wait_until_ready(binding: &str, port: i32) {
let addr = format!("{binding}:{port}");
for _ in 0..200 {
if tokio::net::TcpStream::connect(&addr).await.is_ok() {
return;
}
tokio::time::sleep(tokio::time::Duration::from_millis(25)).await;
}
}
/// Starts the server with a basic route (GET and POST) at the root (`/`), using
/// the given application context.
pub async fn start_from_ctx(ctx: AppContext, port: Option<i32>) -> tokio::task::JoinHandle<()> {
let app_router = AppRoutes::empty()
.add_route(
Routes::new()
.add("/", get(get_action))
.add("/", post(post_action)),
)
.to_router::<AppHook>(ctx.clone(), axum::Router::new())
.expect("to router");
let boot = boot::BootResult {
app_context: ctx,
router: Some(app_router),
worker: None,
run_scheduler: false,
};
start_from_boot(boot, port).await
}
/// Starts the server with a custom route specified by the URI and the HTTP
/// method handler.
pub async fn start_with_route(
ctx: AppContext,
uri: &str,
method: axum::routing::MethodRouter<AppContext>,
port: Option<i32>,
) -> tokio::task::JoinHandle<()> {
let app_router = AppRoutes::empty()
.add_route(Routes::new().add(uri, method))
.to_router::<AppHook>(ctx.clone(), axum::Router::new())
.expect("to router");
let boot = boot::BootResult {
app_context: ctx,
router: Some(app_router),
worker: None,
run_scheduler: false,
};
start_from_boot(boot, port).await
}
+3
View File
@@ -0,0 +1,3 @@
mod build_scripts;
mod controller;
mod infra_cfg;