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,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'",
)