feat(web): Removed Cors impl and replaced with tower-http's CorsLayer (#1123)
Fix #1122, see the issue for more details. Signed-off-by: slinkydeveloper <[email protected]>
This commit is contained in:
@@ -2,6 +2,7 @@ use tonic::{transport::Server, Request, Response, Status};
|
|||||||
|
|
||||||
use hello_world::greeter_server::{Greeter, GreeterServer};
|
use hello_world::greeter_server::{Greeter, GreeterServer};
|
||||||
use hello_world::{HelloReply, HelloRequest};
|
use hello_world::{HelloReply, HelloRequest};
|
||||||
|
use tonic_web::GrpcWebLayer;
|
||||||
|
|
||||||
pub mod hello_world {
|
pub mod hello_world {
|
||||||
tonic::include_proto!("helloworld");
|
tonic::include_proto!("helloworld");
|
||||||
@@ -33,14 +34,12 @@ async fn main() -> Result<(), Box<dyn std::error::Error>> {
|
|||||||
|
|
||||||
let greeter = MyGreeter::default();
|
let greeter = MyGreeter::default();
|
||||||
let greeter = GreeterServer::new(greeter);
|
let greeter = GreeterServer::new(greeter);
|
||||||
let greeter = tonic_web::config()
|
|
||||||
.allow_origins(vec!["127.0.0.1"])
|
|
||||||
.enable(greeter);
|
|
||||||
|
|
||||||
println!("GreeterServer listening on {}", addr);
|
println!("GreeterServer listening on {}", addr);
|
||||||
|
|
||||||
Server::builder()
|
Server::builder()
|
||||||
.accept_http1(true)
|
.accept_http1(true)
|
||||||
|
.layer(GrpcWebLayer::new())
|
||||||
.add_service(greeter)
|
.add_service(greeter)
|
||||||
.serve(addr)
|
.serve(addr)
|
||||||
.await?;
|
.await?;
|
||||||
|
|||||||
@@ -25,6 +25,7 @@ pin-project = "1"
|
|||||||
tonic = {version = "0.8", path = "../tonic", default-features = false, features = ["transport"]}
|
tonic = {version = "0.8", path = "../tonic", default-features = false, features = ["transport"]}
|
||||||
tower-service = "0.3"
|
tower-service = "0.3"
|
||||||
tower-layer = "0.3"
|
tower-layer = "0.3"
|
||||||
|
tower-http = { version = "0.3", features = ["cors"] }
|
||||||
tracing = "0.1"
|
tracing = "0.1"
|
||||||
|
|
||||||
[dev-dependencies]
|
[dev-dependencies]
|
||||||
|
|||||||
@@ -1,166 +0,0 @@
|
|||||||
use std::collections::{BTreeSet, HashSet};
|
|
||||||
use std::convert::TryFrom;
|
|
||||||
use std::time::Duration;
|
|
||||||
|
|
||||||
use http::{header::HeaderName, HeaderValue};
|
|
||||||
use tonic::body::BoxBody;
|
|
||||||
use tower_service::Service;
|
|
||||||
|
|
||||||
use crate::service::GrpcWeb;
|
|
||||||
use crate::BoxError;
|
|
||||||
|
|
||||||
const DEFAULT_MAX_AGE: Duration = Duration::from_secs(24 * 60 * 60);
|
|
||||||
|
|
||||||
const DEFAULT_EXPOSED_HEADERS: [&str; 2] = ["grpc-status", "grpc-message"];
|
|
||||||
|
|
||||||
/// A Configuration builder for grpc_web services.
|
|
||||||
///
|
|
||||||
/// `Config` can be used to tweak the behavior of tonic_web services. Currently,
|
|
||||||
/// `Config` instances only expose cors settings. However, since tonic_web is designed to work
|
|
||||||
/// with grpc-web compliant clients only, some cors options have specific default values and not
|
|
||||||
/// all settings are configurable.
|
|
||||||
///
|
|
||||||
/// ## Default values and configuration options
|
|
||||||
///
|
|
||||||
/// * `allow-origin`: All origins allowed by default. Configurable, but null and wildcard origins
|
|
||||||
/// are not supported.
|
|
||||||
/// * `allow-methods`: `[POST,OPTIONS]`. Not configurable.
|
|
||||||
/// * `allow-headers`: Set to whatever the `OPTIONS` request carries. Not configurable.
|
|
||||||
/// * `allow-credentials`: `true`. Configurable.
|
|
||||||
/// * `max-age`: `86400`. Configurable.
|
|
||||||
/// * `expose-headers`: `grpc-status,grpc-message`. Configurable but values can only be added.
|
|
||||||
/// `grpc-status` and `grpc-message` will always be exposed.
|
|
||||||
#[derive(Debug, Clone)]
|
|
||||||
pub struct Config {
|
|
||||||
pub(crate) allowed_origins: AllowedOrigins,
|
|
||||||
pub(crate) exposed_headers: HashSet<HeaderName>,
|
|
||||||
pub(crate) max_age: Option<Duration>,
|
|
||||||
pub(crate) allow_credentials: bool,
|
|
||||||
}
|
|
||||||
|
|
||||||
#[derive(Debug, Clone)]
|
|
||||||
pub(crate) enum AllowedOrigins {
|
|
||||||
Any,
|
|
||||||
#[allow(clippy::mutable_key_type)]
|
|
||||||
Only(BTreeSet<HeaderValue>),
|
|
||||||
}
|
|
||||||
|
|
||||||
impl AllowedOrigins {
|
|
||||||
pub(crate) fn is_allowed(&self, origin: &HeaderValue) -> bool {
|
|
||||||
match self {
|
|
||||||
AllowedOrigins::Any => true,
|
|
||||||
AllowedOrigins::Only(origins) => origins.contains(origin),
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
impl Config {
|
|
||||||
pub(crate) fn new() -> Config {
|
|
||||||
Config {
|
|
||||||
allowed_origins: AllowedOrigins::Any,
|
|
||||||
exposed_headers: DEFAULT_EXPOSED_HEADERS
|
|
||||||
.iter()
|
|
||||||
.cloned()
|
|
||||||
.map(HeaderName::from_static)
|
|
||||||
.collect(),
|
|
||||||
max_age: Some(DEFAULT_MAX_AGE),
|
|
||||||
allow_credentials: true,
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
/// Allow any origin to access this resource.
|
|
||||||
///
|
|
||||||
/// This is the default value.
|
|
||||||
pub fn allow_all_origins(self) -> Config {
|
|
||||||
Self {
|
|
||||||
allowed_origins: AllowedOrigins::Any,
|
|
||||||
..self
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
/// Only allow a specific set of origins to access this resource.
|
|
||||||
///
|
|
||||||
/// ## Example
|
|
||||||
///
|
|
||||||
/// ```
|
|
||||||
/// tonic_web::config().allow_origins(vec!["http://a.com", "http://b.com"]);
|
|
||||||
/// ```
|
|
||||||
pub fn allow_origins<I>(self, origins: I) -> Config
|
|
||||||
where
|
|
||||||
I: IntoIterator,
|
|
||||||
HeaderValue: TryFrom<I::Item>,
|
|
||||||
{
|
|
||||||
// false positive when using HeaderValue, which uses Bytes internally
|
|
||||||
// https://rust-lang.github.io/rust-clippy/master/index.html#mutable_key_type
|
|
||||||
#[allow(clippy::mutable_key_type)]
|
|
||||||
let origins = origins
|
|
||||||
.into_iter()
|
|
||||||
.map(|v| match TryFrom::try_from(v) {
|
|
||||||
Ok(uri) => uri,
|
|
||||||
Err(_) => panic!("invalid origin"),
|
|
||||||
})
|
|
||||||
.collect();
|
|
||||||
|
|
||||||
Self {
|
|
||||||
allowed_origins: AllowedOrigins::Only(origins),
|
|
||||||
..self
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
/// Adds multiple headers to the list of exposed headers.
|
|
||||||
///
|
|
||||||
/// Default: `grpc-status,grpc-message`. These will always be included.
|
|
||||||
pub fn expose_headers<I>(mut self, headers: I) -> Config
|
|
||||||
where
|
|
||||||
I: IntoIterator,
|
|
||||||
HeaderName: TryFrom<I::Item>,
|
|
||||||
{
|
|
||||||
let iter = headers
|
|
||||||
.into_iter()
|
|
||||||
.map(|header| match TryFrom::try_from(header) {
|
|
||||||
Ok(header) => header,
|
|
||||||
Err(_) => panic!("invalid header"),
|
|
||||||
});
|
|
||||||
|
|
||||||
self.exposed_headers.extend(iter);
|
|
||||||
self
|
|
||||||
}
|
|
||||||
|
|
||||||
/// Defines the maximum cache lifetime for operations allowed on this
|
|
||||||
/// resource.
|
|
||||||
///
|
|
||||||
/// Default: "86400" (24 hours)
|
|
||||||
pub fn max_age<T: Into<Option<Duration>>>(self, max_age: T) -> Config {
|
|
||||||
Self {
|
|
||||||
max_age: max_age.into(),
|
|
||||||
..self
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
/// If true, the `access-control-allow-credentials` will be sent.
|
|
||||||
///
|
|
||||||
/// Default: true
|
|
||||||
pub fn allow_credentials(self, allow_credentials: bool) -> Config {
|
|
||||||
Self {
|
|
||||||
allow_credentials,
|
|
||||||
..self
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
/// enable a tonic service to handle grpc-web requests with this configuration values.
|
|
||||||
pub fn enable<S>(&self, service: S) -> GrpcWeb<S>
|
|
||||||
where
|
|
||||||
S: Service<http::Request<hyper::Body>, Response = http::Response<BoxBody>>,
|
|
||||||
S: Clone + Send + 'static,
|
|
||||||
S::Future: Send + 'static,
|
|
||||||
S::Error: Into<BoxError> + Send,
|
|
||||||
{
|
|
||||||
GrpcWeb::new(service, self.clone())
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
impl Default for Config {
|
|
||||||
fn default() -> Self {
|
|
||||||
Config::new()
|
|
||||||
}
|
|
||||||
}
|
|
||||||
@@ -1,402 +0,0 @@
|
|||||||
use std::sync::Arc;
|
|
||||||
|
|
||||||
pub(crate) use http::header::ACCESS_CONTROL_ALLOW_CREDENTIALS as ALLOW_CREDENTIALS;
|
|
||||||
pub(crate) use http::header::ACCESS_CONTROL_ALLOW_HEADERS as ALLOW_HEADERS;
|
|
||||||
pub(crate) use http::header::ACCESS_CONTROL_ALLOW_METHODS as ALLOW_METHODS;
|
|
||||||
pub(crate) use http::header::ACCESS_CONTROL_ALLOW_ORIGIN as ALLOW_ORIGIN;
|
|
||||||
pub(crate) use http::header::ACCESS_CONTROL_EXPOSE_HEADERS as EXPOSE_HEADERS;
|
|
||||||
pub(crate) use http::header::ACCESS_CONTROL_MAX_AGE as MAX_AGE;
|
|
||||||
pub(crate) use http::header::ACCESS_CONTROL_REQUEST_HEADERS as REQUEST_HEADERS;
|
|
||||||
pub(crate) use http::header::ACCESS_CONTROL_REQUEST_METHOD as REQUEST_METHOD;
|
|
||||||
pub(crate) use http::header::ORIGIN;
|
|
||||||
use http::{header, HeaderMap, HeaderValue, Method};
|
|
||||||
use tracing::debug;
|
|
||||||
|
|
||||||
use crate::config::Config;
|
|
||||||
|
|
||||||
const DEFAULT_ALLOWED_METHODS: &[Method; 2] = &[Method::POST, Method::OPTIONS];
|
|
||||||
|
|
||||||
#[derive(Debug, Clone)]
|
|
||||||
pub(crate) struct Cors {
|
|
||||||
cache: Arc<Cache>,
|
|
||||||
}
|
|
||||||
|
|
||||||
#[derive(Debug, PartialEq)]
|
|
||||||
pub(crate) enum Error {
|
|
||||||
OriginNotAllowed,
|
|
||||||
MethodNotAllowed,
|
|
||||||
}
|
|
||||||
|
|
||||||
#[derive(Clone, Debug)]
|
|
||||||
struct Cache {
|
|
||||||
config: Config,
|
|
||||||
expose_headers: HeaderValue,
|
|
||||||
allow_methods: HeaderValue,
|
|
||||||
allow_credentials: HeaderValue,
|
|
||||||
}
|
|
||||||
|
|
||||||
impl Cors {
|
|
||||||
pub(crate) fn new(config: Config) -> Cors {
|
|
||||||
let expose_headers = join_header_value(&config.exposed_headers).unwrap();
|
|
||||||
let allow_methods = HeaderValue::from_static("POST,OPTIONS");
|
|
||||||
let allow_credentials = HeaderValue::from_static("true");
|
|
||||||
|
|
||||||
let cache = Arc::new(Cache {
|
|
||||||
config,
|
|
||||||
expose_headers,
|
|
||||||
allow_methods,
|
|
||||||
allow_credentials,
|
|
||||||
});
|
|
||||||
|
|
||||||
Cors { cache }
|
|
||||||
}
|
|
||||||
|
|
||||||
fn is_method_allowed(&self, header: Option<&HeaderValue>) -> bool {
|
|
||||||
match header {
|
|
||||||
Some(value) => match Method::from_bytes(value.as_bytes()) {
|
|
||||||
Ok(method) => DEFAULT_ALLOWED_METHODS.contains(&method),
|
|
||||||
Err(_) => {
|
|
||||||
debug!("access-control-request-method {:?} is not valid", value);
|
|
||||||
false
|
|
||||||
}
|
|
||||||
},
|
|
||||||
None => {
|
|
||||||
debug!("access-control-request-method is missing");
|
|
||||||
false
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
pub(crate) fn preflight(
|
|
||||||
&self,
|
|
||||||
req_headers: &HeaderMap,
|
|
||||||
origin: &HeaderValue,
|
|
||||||
request_headers_header: &HeaderValue,
|
|
||||||
) -> Result<HeaderMap, Error> {
|
|
||||||
if !self.is_origin_allowed(origin) {
|
|
||||||
return Err(Error::OriginNotAllowed);
|
|
||||||
}
|
|
||||||
|
|
||||||
if !self.is_method_allowed(req_headers.get(REQUEST_METHOD)) {
|
|
||||||
return Err(Error::MethodNotAllowed);
|
|
||||||
}
|
|
||||||
|
|
||||||
let mut headers = self.common_headers(origin.clone());
|
|
||||||
headers.insert(ALLOW_METHODS, self.cache.allow_methods.clone());
|
|
||||||
headers.insert(ALLOW_HEADERS, request_headers_header.clone());
|
|
||||||
|
|
||||||
if let Some(max_age) = self.cache.config.max_age {
|
|
||||||
headers.insert(MAX_AGE, HeaderValue::from(max_age.as_secs()));
|
|
||||||
}
|
|
||||||
|
|
||||||
Ok(headers)
|
|
||||||
}
|
|
||||||
|
|
||||||
pub(crate) fn simple(&self, headers: &HeaderMap) -> Result<HeaderMap, Error> {
|
|
||||||
match headers.get(header::ORIGIN) {
|
|
||||||
Some(origin) if self.is_origin_allowed(origin) => {
|
|
||||||
Ok(self.common_headers(origin.clone()))
|
|
||||||
}
|
|
||||||
Some(_) => Err(Error::OriginNotAllowed),
|
|
||||||
None => Ok(HeaderMap::new()),
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
fn common_headers(&self, origin: HeaderValue) -> HeaderMap {
|
|
||||||
let mut headers = HeaderMap::new();
|
|
||||||
headers.insert(ALLOW_ORIGIN, origin);
|
|
||||||
headers.insert(EXPOSE_HEADERS, self.cache.expose_headers.clone());
|
|
||||||
|
|
||||||
if self.cache.config.allow_credentials {
|
|
||||||
headers.insert(ALLOW_CREDENTIALS, self.cache.allow_credentials.clone());
|
|
||||||
}
|
|
||||||
|
|
||||||
headers
|
|
||||||
}
|
|
||||||
|
|
||||||
fn is_origin_allowed(&self, origin: &HeaderValue) -> bool {
|
|
||||||
self.cache.config.allowed_origins.is_allowed(origin)
|
|
||||||
}
|
|
||||||
|
|
||||||
#[cfg(test)]
|
|
||||||
pub(crate) fn __check_preflight(&self, headers: &HeaderMap) -> Result<HeaderMap, Error> {
|
|
||||||
self.preflight(
|
|
||||||
headers,
|
|
||||||
headers.get(ORIGIN).unwrap(),
|
|
||||||
headers.get(REQUEST_HEADERS).unwrap(),
|
|
||||||
)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
#[cfg(test)]
|
|
||||||
impl Default for Cors {
|
|
||||||
fn default() -> Self {
|
|
||||||
Cors::new(Config::default())
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
fn join_header_value<I>(values: I) -> Result<HeaderValue, header::InvalidHeaderValue>
|
|
||||||
where
|
|
||||||
I: IntoIterator,
|
|
||||||
I::Item: AsRef<str>,
|
|
||||||
{
|
|
||||||
let mut values = values.into_iter();
|
|
||||||
let mut value = Vec::new();
|
|
||||||
|
|
||||||
if let Some(v) = values.next() {
|
|
||||||
value.extend(v.as_ref().as_bytes());
|
|
||||||
}
|
|
||||||
for v in values {
|
|
||||||
value.push(b',');
|
|
||||||
value.extend(v.as_ref().as_bytes());
|
|
||||||
}
|
|
||||||
HeaderValue::from_bytes(&value)
|
|
||||||
}
|
|
||||||
|
|
||||||
#[cfg(test)]
|
|
||||||
mod tests {
|
|
||||||
use super::*;
|
|
||||||
|
|
||||||
macro_rules! assert_value_eq {
|
|
||||||
($header:expr, $expected:expr) => {
|
|
||||||
fn sorted(value: &str) -> Vec<&str> {
|
|
||||||
let mut vec = value.split(",").collect::<Vec<_>>();
|
|
||||||
vec.sort_unstable();
|
|
||||||
vec
|
|
||||||
}
|
|
||||||
|
|
||||||
assert_eq!(sorted($header.to_str().unwrap()), sorted($expected))
|
|
||||||
};
|
|
||||||
}
|
|
||||||
|
|
||||||
fn value(s: &str) -> HeaderValue {
|
|
||||||
s.parse().unwrap()
|
|
||||||
}
|
|
||||||
|
|
||||||
impl From<Config> for Cors {
|
|
||||||
fn from(c: Config) -> Self {
|
|
||||||
Cors::new(c)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
#[should_panic]
|
|
||||||
#[ignore]
|
|
||||||
fn origin_is_valid_url() {
|
|
||||||
Config::new().allow_origins(vec!["foo"]);
|
|
||||||
}
|
|
||||||
|
|
||||||
mod preflight {
|
|
||||||
use super::*;
|
|
||||||
|
|
||||||
fn preflight_headers() -> HeaderMap {
|
|
||||||
let mut headers = HeaderMap::new();
|
|
||||||
headers.insert(ORIGIN, value("http://example.com"));
|
|
||||||
headers.insert(REQUEST_METHOD, value("POST"));
|
|
||||||
headers.insert(REQUEST_HEADERS, value("x-grpc-web"));
|
|
||||||
headers
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn default_config() {
|
|
||||||
let cors = Cors::default();
|
|
||||||
let headers = cors.__check_preflight(&preflight_headers()).unwrap();
|
|
||||||
|
|
||||||
assert_eq!(headers[ALLOW_ORIGIN], "http://example.com");
|
|
||||||
assert_eq!(headers[ALLOW_METHODS], "POST,OPTIONS");
|
|
||||||
assert_eq!(headers[ALLOW_HEADERS], "x-grpc-web");
|
|
||||||
assert_eq!(headers[ALLOW_CREDENTIALS], "true");
|
|
||||||
assert_eq!(headers[MAX_AGE], "86400");
|
|
||||||
assert_value_eq!(&headers[EXPOSE_HEADERS], "grpc-status,grpc-message");
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn any_origin() {
|
|
||||||
let cors: Cors = Config::new().allow_all_origins().into();
|
|
||||||
|
|
||||||
assert!(cors.__check_preflight(&preflight_headers()).is_ok());
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn origin_list() {
|
|
||||||
let cors: Cors = Config::new()
|
|
||||||
.allow_origins(vec![
|
|
||||||
HeaderValue::from_static("http://a.com"),
|
|
||||||
HeaderValue::from_static("http://b.com"),
|
|
||||||
])
|
|
||||||
.into();
|
|
||||||
|
|
||||||
let mut req_headers = preflight_headers();
|
|
||||||
req_headers.insert(ORIGIN, value("http://b.com"));
|
|
||||||
|
|
||||||
assert!(cors.__check_preflight(&req_headers).is_ok());
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn origin_not_allowed() {
|
|
||||||
let cors: Cors = Config::new().allow_origins(vec!["http://a.com"]).into();
|
|
||||||
|
|
||||||
let err = cors.__check_preflight(&preflight_headers()).unwrap_err();
|
|
||||||
|
|
||||||
assert_eq!(err, Error::OriginNotAllowed)
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn disallow_credentials() {
|
|
||||||
let cors = Cors::new(Config::new().allow_credentials(false));
|
|
||||||
let headers = cors.__check_preflight(&preflight_headers()).unwrap();
|
|
||||||
|
|
||||||
assert!(!headers.contains_key(ALLOW_CREDENTIALS));
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn expose_headers_are_merged() {
|
|
||||||
let cors = Cors::new(Config::new().expose_headers(vec!["x-request-id"]));
|
|
||||||
let headers = cors.__check_preflight(&preflight_headers()).unwrap();
|
|
||||||
|
|
||||||
assert_value_eq!(
|
|
||||||
&headers[EXPOSE_HEADERS],
|
|
||||||
"x-request-id,grpc-message,grpc-status"
|
|
||||||
);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn allow_headers_echo_request_headers() {
|
|
||||||
let cors = Cors::default();
|
|
||||||
let mut request_headers = preflight_headers();
|
|
||||||
request_headers.insert(REQUEST_HEADERS, value("x-grpc-web,foo,x-request-id"));
|
|
||||||
|
|
||||||
let headers = cors.__check_preflight(&request_headers).unwrap();
|
|
||||||
|
|
||||||
assert_value_eq!(&headers[ALLOW_HEADERS], "x-grpc-web,foo,x-request-id");
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn missing_request_method() {
|
|
||||||
let cors = Cors::default();
|
|
||||||
let mut request_headers = preflight_headers();
|
|
||||||
request_headers.remove(REQUEST_METHOD);
|
|
||||||
|
|
||||||
let err = cors.__check_preflight(&request_headers).unwrap_err();
|
|
||||||
|
|
||||||
assert_eq!(err, Error::MethodNotAllowed);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn only_options_and_post_allowed() {
|
|
||||||
let cors = Cors::default();
|
|
||||||
|
|
||||||
for method in &[
|
|
||||||
Method::GET,
|
|
||||||
Method::DELETE,
|
|
||||||
Method::TRACE,
|
|
||||||
Method::PATCH,
|
|
||||||
Method::PUT,
|
|
||||||
Method::HEAD,
|
|
||||||
] {
|
|
||||||
let mut request_headers = preflight_headers();
|
|
||||||
request_headers.insert(REQUEST_METHOD, value(method.as_str()));
|
|
||||||
|
|
||||||
assert_eq!(
|
|
||||||
cors.__check_preflight(&request_headers).unwrap_err(),
|
|
||||||
Error::MethodNotAllowed,
|
|
||||||
)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn custom_max_age() {
|
|
||||||
use std::time::Duration;
|
|
||||||
|
|
||||||
let cors = Cors::new(Config::new().max_age(Duration::from_secs(99)));
|
|
||||||
let headers = cors.__check_preflight(&preflight_headers()).unwrap();
|
|
||||||
|
|
||||||
assert_eq!(headers[MAX_AGE], "99");
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn no_max_age() {
|
|
||||||
let cors = Cors::new(Config::new().max_age(None));
|
|
||||||
let headers = cors.__check_preflight(&preflight_headers()).unwrap();
|
|
||||||
|
|
||||||
assert!(!headers.contains_key(MAX_AGE));
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
mod simple {
|
|
||||||
use super::*;
|
|
||||||
|
|
||||||
fn request_headers() -> HeaderMap {
|
|
||||||
let mut headers = HeaderMap::new();
|
|
||||||
headers.insert(ORIGIN, value("http://example.com"));
|
|
||||||
headers
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn default_config() {
|
|
||||||
let cors = Cors::default();
|
|
||||||
let headers = cors.simple(&request_headers()).unwrap();
|
|
||||||
|
|
||||||
assert_eq!(headers[ALLOW_ORIGIN], "http://example.com");
|
|
||||||
assert_eq!(headers[ALLOW_CREDENTIALS], "true");
|
|
||||||
assert_value_eq!(&headers[EXPOSE_HEADERS], "grpc-message,grpc-status");
|
|
||||||
|
|
||||||
assert!(!headers.contains_key(ALLOW_HEADERS));
|
|
||||||
assert!(!headers.contains_key(ALLOW_METHODS));
|
|
||||||
assert!(!headers.contains_key(MAX_AGE));
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn any_origin() {
|
|
||||||
let cors: Cors = Config::new().allow_all_origins().into();
|
|
||||||
|
|
||||||
assert!(cors.simple(&request_headers()).is_ok());
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn origin_list() {
|
|
||||||
let cors: Cors = Config::new()
|
|
||||||
.allow_origins(vec![
|
|
||||||
HeaderValue::from_static("http://a.com"),
|
|
||||||
HeaderValue::from_static("http://b.com"),
|
|
||||||
])
|
|
||||||
.into();
|
|
||||||
|
|
||||||
let mut req_headers = request_headers();
|
|
||||||
req_headers.insert(ORIGIN, value("http://b.com"));
|
|
||||||
|
|
||||||
assert!(cors.simple(&req_headers).is_ok());
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn origin_not_allowed() {
|
|
||||||
let cors: Cors = Config::new().allow_origins(vec!["http://a.com"]).into();
|
|
||||||
|
|
||||||
let err = cors.simple(&request_headers()).unwrap_err();
|
|
||||||
|
|
||||||
assert_eq!(err, Error::OriginNotAllowed)
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn disallow_credentials() {
|
|
||||||
let cors = Cors::new(Config::new().allow_credentials(false));
|
|
||||||
let headers = cors.simple(&request_headers()).unwrap();
|
|
||||||
|
|
||||||
assert!(!headers.contains_key(ALLOW_CREDENTIALS));
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn expose_headers_are_merged() {
|
|
||||||
let cors: Cors = Config::new()
|
|
||||||
.expose_headers(vec!["x-hello", "custom-1"])
|
|
||||||
.into();
|
|
||||||
|
|
||||||
let headers = cors.simple(&request_headers()).unwrap();
|
|
||||||
|
|
||||||
assert_value_eq!(
|
|
||||||
&headers[EXPOSE_HEADERS],
|
|
||||||
"grpc-message,grpc-status,x-hello,custom-1"
|
|
||||||
);
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
@@ -1,4 +1,4 @@
|
|||||||
use super::{BoxBody, BoxError, Config, GrpcWeb};
|
use super::{BoxBody, BoxError, GrpcWebService};
|
||||||
|
|
||||||
use tower_layer::Layer;
|
use tower_layer::Layer;
|
||||||
use tower_service::Service;
|
use tower_service::Service;
|
||||||
@@ -23,9 +23,9 @@ where
|
|||||||
S::Future: Send + 'static,
|
S::Future: Send + 'static,
|
||||||
S::Error: Into<BoxError> + Send,
|
S::Error: Into<BoxError> + Send,
|
||||||
{
|
{
|
||||||
type Service = GrpcWeb<S>;
|
type Service = GrpcWebService<S>;
|
||||||
|
|
||||||
fn layer(&self, inner: S) -> Self::Service {
|
fn layer(&self, inner: S) -> Self::Service {
|
||||||
Config::default().enable(inner)
|
GrpcWebService::new(inner)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
+37
-32
@@ -34,7 +34,8 @@
|
|||||||
//!
|
//!
|
||||||
//! ```
|
//! ```
|
||||||
//! This will apply a default configuration that works well with grpc-web clients out of the box.
|
//! This will apply a default configuration that works well with grpc-web clients out of the box.
|
||||||
//! See the [`Config`] documentation for details.
|
//!
|
||||||
|
//! You can customize the CORS configuration composing the [`GrpcWebLayer`] with the cors layer of your choice.
|
||||||
//!
|
//!
|
||||||
//! Alternatively, if you have a tls enabled server, you could skip setting `accept_http1` to `true`.
|
//! Alternatively, if you have a tls enabled server, you could skip setting `accept_http1` to `true`.
|
||||||
//! This works because the browser will handle `ALPN`.
|
//! This works because the browser will handle `ALPN`.
|
||||||
@@ -77,7 +78,6 @@
|
|||||||
//! [grpc-web]: https://github.com/grpc/grpc-web
|
//! [grpc-web]: https://github.com/grpc/grpc-web
|
||||||
//! [tower]: https://github.com/tower-rs/tower
|
//! [tower]: https://github.com/tower-rs/tower
|
||||||
//! [`enable`]: crate::enable()
|
//! [`enable`]: crate::enable()
|
||||||
//! [`Config`]: crate::Config
|
|
||||||
#![warn(
|
#![warn(
|
||||||
missing_debug_implementations,
|
missing_debug_implementations,
|
||||||
missing_docs,
|
missing_docs,
|
||||||
@@ -87,50 +87,55 @@
|
|||||||
#![doc(html_root_url = "https://docs.rs/tonic-web/0.4.0")]
|
#![doc(html_root_url = "https://docs.rs/tonic-web/0.4.0")]
|
||||||
#![doc(issue_tracker_base_url = "https://github.com/hyperium/tonic/issues/")]
|
#![doc(issue_tracker_base_url = "https://github.com/hyperium/tonic/issues/")]
|
||||||
|
|
||||||
pub use config::Config;
|
|
||||||
pub use layer::GrpcWebLayer;
|
pub use layer::GrpcWebLayer;
|
||||||
pub use service::GrpcWeb;
|
pub use service::{GrpcWebService, ResponseFuture};
|
||||||
|
|
||||||
mod call;
|
mod call;
|
||||||
mod config;
|
|
||||||
mod cors;
|
|
||||||
mod layer;
|
mod layer;
|
||||||
mod service;
|
mod service;
|
||||||
|
|
||||||
use std::future::Future;
|
use http::header::HeaderName;
|
||||||
use std::pin::Pin;
|
use std::time::Duration;
|
||||||
use tonic::body::BoxBody;
|
use tonic::body::BoxBody;
|
||||||
|
use tower_http::cors::{AllowOrigin, Cors, CorsLayer};
|
||||||
|
use tower_layer::Layer;
|
||||||
use tower_service::Service;
|
use tower_service::Service;
|
||||||
|
|
||||||
/// enable a tonic service to handle grpc-web requests with the default configuration.
|
const DEFAULT_MAX_AGE: Duration = Duration::from_secs(24 * 60 * 60);
|
||||||
|
const DEFAULT_EXPOSED_HEADERS: [&str; 3] =
|
||||||
|
["grpc-status", "grpc-message", "grpc-status-details-bin"];
|
||||||
|
const DEFAULT_ALLOW_HEADERS: [&str; 4] =
|
||||||
|
["x-grpc-web", "content-type", "x-user-agent", "grpc-timeout"];
|
||||||
|
|
||||||
|
type BoxError = Box<dyn std::error::Error + Send + Sync>;
|
||||||
|
|
||||||
|
/// Enable a tonic service to handle grpc-web requests with the default configuration.
|
||||||
///
|
///
|
||||||
/// Shortcut for `tonic_web::config().enable(service)`
|
/// You can customize the CORS configuration composing the [`GrpcWebLayer`] with the cors layer of your choice.
|
||||||
pub fn enable<S>(service: S) -> GrpcWeb<S>
|
pub fn enable<S>(service: S) -> Cors<GrpcWebService<S>>
|
||||||
where
|
where
|
||||||
S: Service<http::Request<hyper::Body>, Response = http::Response<BoxBody>>,
|
S: Service<http::Request<hyper::Body>, Response = http::Response<BoxBody>>,
|
||||||
S: Clone + Send + 'static,
|
S: Clone + Send + 'static,
|
||||||
S::Future: Send + 'static,
|
S::Future: Send + 'static,
|
||||||
S::Error: Into<BoxError> + Send,
|
S::Error: Into<BoxError> + Send,
|
||||||
{
|
{
|
||||||
config().enable(service)
|
CorsLayer::new()
|
||||||
|
.allow_origin(AllowOrigin::mirror_request())
|
||||||
|
.allow_credentials(true)
|
||||||
|
.max_age(DEFAULT_MAX_AGE)
|
||||||
|
.expose_headers(
|
||||||
|
DEFAULT_EXPOSED_HEADERS
|
||||||
|
.iter()
|
||||||
|
.cloned()
|
||||||
|
.map(HeaderName::from_static)
|
||||||
|
.collect::<Vec<HeaderName>>(),
|
||||||
|
)
|
||||||
|
.allow_headers(
|
||||||
|
DEFAULT_ALLOW_HEADERS
|
||||||
|
.iter()
|
||||||
|
.cloned()
|
||||||
|
.map(HeaderName::from_static)
|
||||||
|
.collect::<Vec<HeaderName>>(),
|
||||||
|
)
|
||||||
|
.layer(GrpcWebService::new(service))
|
||||||
}
|
}
|
||||||
|
|
||||||
/// returns a default [`Config`] instance for configuring services.
|
|
||||||
///
|
|
||||||
/// ## Example
|
|
||||||
///
|
|
||||||
/// ```
|
|
||||||
/// let config = tonic_web::config()
|
|
||||||
/// .allow_origins(vec!["http://foo.com"])
|
|
||||||
/// .allow_credentials(false)
|
|
||||||
/// .expose_headers(vec!["x-request-id"]);
|
|
||||||
///
|
|
||||||
/// // let greeter = config.enable(Greeter);
|
|
||||||
/// // let route_guide = config.enable(RouteGuide);
|
|
||||||
/// ```
|
|
||||||
pub fn config() -> Config {
|
|
||||||
Config::default()
|
|
||||||
}
|
|
||||||
|
|
||||||
type BoxError = Box<dyn std::error::Error + Send + Sync>;
|
|
||||||
type BoxFuture<T, E> = Pin<Box<dyn Future<Output = Result<T, E>> + Send>>;
|
|
||||||
|
|||||||
+101
-218
@@ -1,7 +1,11 @@
|
|||||||
|
use futures_core::ready;
|
||||||
|
use std::future::Future;
|
||||||
|
use std::pin::Pin;
|
||||||
use std::task::{Context, Poll};
|
use std::task::{Context, Poll};
|
||||||
|
|
||||||
use http::{header, HeaderMap, HeaderValue, Method, Request, Response, StatusCode, Version};
|
use http::{header, HeaderMap, HeaderValue, Method, Request, Response, StatusCode, Version};
|
||||||
use hyper::Body;
|
use hyper::Body;
|
||||||
|
use pin_project::pin_project;
|
||||||
use tonic::body::{empty_body, BoxBody};
|
use tonic::body::{empty_body, BoxBody};
|
||||||
use tonic::transport::NamedService;
|
use tonic::transport::NamedService;
|
||||||
use tower_service::Service;
|
use tower_service::Service;
|
||||||
@@ -9,17 +13,14 @@ use tracing::{debug, trace};
|
|||||||
|
|
||||||
use crate::call::content_types::is_grpc_web;
|
use crate::call::content_types::is_grpc_web;
|
||||||
use crate::call::{Encoding, GrpcWebCall};
|
use crate::call::{Encoding, GrpcWebCall};
|
||||||
use crate::cors::Cors;
|
use crate::BoxError;
|
||||||
use crate::cors::{ORIGIN, REQUEST_HEADERS};
|
|
||||||
use crate::{BoxError, BoxFuture, Config};
|
|
||||||
|
|
||||||
const GRPC: &str = "application/grpc";
|
const GRPC: &str = "application/grpc";
|
||||||
|
|
||||||
/// Service implementing the grpc-web protocol.
|
/// Service implementing the grpc-web protocol.
|
||||||
#[derive(Debug, Clone)]
|
#[derive(Debug, Clone)]
|
||||||
pub struct GrpcWeb<S> {
|
pub struct GrpcWebService<S> {
|
||||||
inner: S,
|
inner: S,
|
||||||
cors: Cors,
|
|
||||||
}
|
}
|
||||||
|
|
||||||
#[derive(Debug, PartialEq)]
|
#[derive(Debug, PartialEq)]
|
||||||
@@ -36,55 +37,35 @@ enum RequestKind<'a> {
|
|||||||
encoding: Encoding,
|
encoding: Encoding,
|
||||||
accept: Encoding,
|
accept: Encoding,
|
||||||
},
|
},
|
||||||
// The request is considered a grpc-web preflight request if all these
|
|
||||||
// conditions are met:
|
|
||||||
//
|
|
||||||
// - the request method is `OPTIONS`
|
|
||||||
// - request headers include `origin`
|
|
||||||
// - `access-control-request-headers` header is present and includes `x-grpc-web`
|
|
||||||
GrpcWebPreflight {
|
|
||||||
origin: &'a HeaderValue,
|
|
||||||
request_headers: &'a HeaderValue,
|
|
||||||
},
|
|
||||||
// All other requests, including `application/grpc`
|
// All other requests, including `application/grpc`
|
||||||
Other(http::Version),
|
Other(http::Version),
|
||||||
}
|
}
|
||||||
|
|
||||||
impl<S> GrpcWeb<S> {
|
impl<S> GrpcWebService<S> {
|
||||||
pub(crate) fn new(inner: S, config: Config) -> Self {
|
pub(crate) fn new(inner: S) -> Self {
|
||||||
GrpcWeb {
|
GrpcWebService { inner }
|
||||||
inner,
|
}
|
||||||
cors: Cors::new(config),
|
}
|
||||||
|
|
||||||
|
impl<S> GrpcWebService<S>
|
||||||
|
where
|
||||||
|
S: Service<Request<Body>, Response = Response<BoxBody>> + Send + 'static,
|
||||||
|
{
|
||||||
|
fn response(&self, status: StatusCode) -> ResponseFuture<S::Future> {
|
||||||
|
ResponseFuture {
|
||||||
|
case: Case::ImmediateResponse {
|
||||||
|
res: Some(
|
||||||
|
Response::builder()
|
||||||
|
.status(status)
|
||||||
|
.body(empty_body())
|
||||||
|
.unwrap(),
|
||||||
|
),
|
||||||
|
},
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
impl<S> GrpcWeb<S>
|
impl<S> Service<Request<Body>> for GrpcWebService<S>
|
||||||
where
|
|
||||||
S: Service<Request<Body>, Response = Response<BoxBody>> + Send + 'static,
|
|
||||||
{
|
|
||||||
fn no_content(&self, headers: HeaderMap) -> BoxFuture<S::Response, S::Error> {
|
|
||||||
let mut res = Response::builder()
|
|
||||||
.status(StatusCode::NO_CONTENT)
|
|
||||||
.body(empty_body())
|
|
||||||
.unwrap();
|
|
||||||
|
|
||||||
res.headers_mut().extend(headers);
|
|
||||||
|
|
||||||
Box::pin(async { Ok(res) })
|
|
||||||
}
|
|
||||||
|
|
||||||
fn response(&self, status: StatusCode) -> BoxFuture<S::Response, S::Error> {
|
|
||||||
Box::pin(async move {
|
|
||||||
Ok(Response::builder()
|
|
||||||
.status(status)
|
|
||||||
.body(empty_body())
|
|
||||||
.unwrap())
|
|
||||||
})
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
impl<S> Service<Request<Body>> for GrpcWeb<S>
|
|
||||||
where
|
where
|
||||||
S: Service<Request<Body>, Response = Response<BoxBody>> + Send + 'static,
|
S: Service<Request<Body>, Response = Response<BoxBody>> + Send + 'static,
|
||||||
S::Future: Send + 'static,
|
S::Future: Send + 'static,
|
||||||
@@ -92,7 +73,7 @@ where
|
|||||||
{
|
{
|
||||||
type Response = S::Response;
|
type Response = S::Response;
|
||||||
type Error = S::Error;
|
type Error = S::Error;
|
||||||
type Future = BoxFuture<Self::Response, Self::Error>;
|
type Future = ResponseFuture<S::Future>;
|
||||||
|
|
||||||
fn poll_ready(&mut self, cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
|
fn poll_ready(&mut self, cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
|
||||||
self.inner.poll_ready(cx)
|
self.inner.poll_ready(cx)
|
||||||
@@ -113,23 +94,16 @@ where
|
|||||||
method: &Method::POST,
|
method: &Method::POST,
|
||||||
encoding,
|
encoding,
|
||||||
accept,
|
accept,
|
||||||
} => match self.cors.simple(req.headers()) {
|
} => {
|
||||||
Ok(headers) => {
|
trace!(kind = "simple", path = ?req.uri().path(), ?encoding, ?accept);
|
||||||
trace!(kind = "simple", path = ?req.uri().path(), ?encoding, ?accept);
|
|
||||||
|
|
||||||
let fut = self.inner.call(coerce_request(req, encoding));
|
ResponseFuture {
|
||||||
|
case: Case::GrpcWeb {
|
||||||
Box::pin(async move {
|
future: self.inner.call(coerce_request(req, encoding)),
|
||||||
let mut res = coerce_response(fut.await?, accept);
|
accept,
|
||||||
res.headers_mut().extend(headers);
|
},
|
||||||
Ok(res)
|
|
||||||
})
|
|
||||||
}
|
}
|
||||||
Err(e) => {
|
}
|
||||||
debug!(kind = "simple", error=?e, ?req);
|
|
||||||
self.response(StatusCode::FORBIDDEN)
|
|
||||||
}
|
|
||||||
},
|
|
||||||
|
|
||||||
// The request's content-type matches one of the 4 supported grpc-web
|
// The request's content-type matches one of the 4 supported grpc-web
|
||||||
// content-types, but the request method is not `POST`.
|
// content-types, but the request method is not `POST`.
|
||||||
@@ -139,27 +113,15 @@ where
|
|||||||
self.response(StatusCode::METHOD_NOT_ALLOWED)
|
self.response(StatusCode::METHOD_NOT_ALLOWED)
|
||||||
}
|
}
|
||||||
|
|
||||||
// A valid grpc-web preflight request, regardless of HTTP version.
|
// All http/2 requests that are not grpc-web are passed through to the inner service,
|
||||||
// This is handled by the cors module.
|
// whatever they are.
|
||||||
RequestKind::GrpcWebPreflight {
|
|
||||||
origin,
|
|
||||||
request_headers,
|
|
||||||
} => match self.cors.preflight(req.headers(), origin, request_headers) {
|
|
||||||
Ok(headers) => {
|
|
||||||
trace!(kind = "preflight", path = ?req.uri().path(), ?origin);
|
|
||||||
self.no_content(headers)
|
|
||||||
}
|
|
||||||
Err(e) => {
|
|
||||||
debug!(kind = "preflight", error = ?e, ?req);
|
|
||||||
self.response(StatusCode::FORBIDDEN)
|
|
||||||
}
|
|
||||||
},
|
|
||||||
|
|
||||||
// All http/2 requests that are not grpc-web or grpc-web preflight
|
|
||||||
// are passed through to the inner service, whatever they are.
|
|
||||||
RequestKind::Other(Version::HTTP_2) => {
|
RequestKind::Other(Version::HTTP_2) => {
|
||||||
debug!(kind = "other h2", content_type = ?req.headers().get(header::CONTENT_TYPE));
|
debug!(kind = "other h2", content_type = ?req.headers().get(header::CONTENT_TYPE));
|
||||||
Box::pin(self.inner.call(req))
|
ResponseFuture {
|
||||||
|
case: Case::Other {
|
||||||
|
future: self.inner.call(req),
|
||||||
|
},
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// Return HTTP 400 for all other requests.
|
// Return HTTP 400 for all other requests.
|
||||||
@@ -171,7 +133,54 @@ where
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
impl<S: NamedService> NamedService for GrpcWeb<S> {
|
/// Response future for the [`GrpcWebService`].
|
||||||
|
#[allow(missing_debug_implementations)]
|
||||||
|
#[pin_project]
|
||||||
|
#[must_use = "futures do nothing unless polled"]
|
||||||
|
pub struct ResponseFuture<F> {
|
||||||
|
#[pin]
|
||||||
|
case: Case<F>,
|
||||||
|
}
|
||||||
|
|
||||||
|
#[pin_project(project = CaseProj)]
|
||||||
|
enum Case<F> {
|
||||||
|
GrpcWeb {
|
||||||
|
#[pin]
|
||||||
|
future: F,
|
||||||
|
accept: Encoding,
|
||||||
|
},
|
||||||
|
Other {
|
||||||
|
#[pin]
|
||||||
|
future: F,
|
||||||
|
},
|
||||||
|
ImmediateResponse {
|
||||||
|
res: Option<Response<BoxBody>>,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
impl<F, E> Future for ResponseFuture<F>
|
||||||
|
where
|
||||||
|
F: Future<Output = Result<Response<BoxBody>, E>> + Send + 'static,
|
||||||
|
E: Into<BoxError> + Send,
|
||||||
|
{
|
||||||
|
type Output = Result<Response<BoxBody>, E>;
|
||||||
|
|
||||||
|
fn poll(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> {
|
||||||
|
let mut this = self.project();
|
||||||
|
|
||||||
|
match this.case.as_mut().project() {
|
||||||
|
CaseProj::GrpcWeb { future, accept } => {
|
||||||
|
let res = ready!(future.poll(cx))?;
|
||||||
|
|
||||||
|
Poll::Ready(Ok(coerce_response(res, *accept)))
|
||||||
|
}
|
||||||
|
CaseProj::Other { future } => future.poll(cx),
|
||||||
|
CaseProj::ImmediateResponse { res } => Poll::Ready(Ok(res.take().unwrap())),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
impl<S: NamedService> NamedService for GrpcWebService<S> {
|
||||||
const NAME: &'static str = S::NAME;
|
const NAME: &'static str = S::NAME;
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -185,20 +194,6 @@ impl<'a> RequestKind<'a> {
|
|||||||
};
|
};
|
||||||
}
|
}
|
||||||
|
|
||||||
if let (&Method::OPTIONS, Some(origin), Some(value)) =
|
|
||||||
(method, headers.get(ORIGIN), headers.get(REQUEST_HEADERS))
|
|
||||||
{
|
|
||||||
match value.to_str() {
|
|
||||||
Ok(h) if h.contains("x-grpc-web") => {
|
|
||||||
return RequestKind::GrpcWebPreflight {
|
|
||||||
origin,
|
|
||||||
request_headers: value,
|
|
||||||
};
|
|
||||||
}
|
|
||||||
_ => {}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
RequestKind::Other(version)
|
RequestKind::Other(version)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -241,9 +236,13 @@ fn coerce_response(res: Response<BoxBody>, encoding: Encoding) -> Response<BoxBo
|
|||||||
mod tests {
|
mod tests {
|
||||||
use super::*;
|
use super::*;
|
||||||
use crate::call::content_types::*;
|
use crate::call::content_types::*;
|
||||||
use http::header::{CONTENT_TYPE, ORIGIN};
|
use http::header::{
|
||||||
|
ACCESS_CONTROL_REQUEST_HEADERS, ACCESS_CONTROL_REQUEST_METHOD, CONTENT_TYPE, ORIGIN,
|
||||||
|
};
|
||||||
|
|
||||||
#[derive(Clone)]
|
type BoxFuture<T, E> = Pin<Box<dyn Future<Output = Result<T, E>> + Send>>;
|
||||||
|
|
||||||
|
#[derive(Debug, Clone)]
|
||||||
struct Svc;
|
struct Svc;
|
||||||
|
|
||||||
impl tower_service::Service<Request<Body>> for Svc {
|
impl tower_service::Service<Request<Body>> for Svc {
|
||||||
@@ -307,18 +306,7 @@ mod tests {
|
|||||||
}
|
}
|
||||||
|
|
||||||
#[tokio::test]
|
#[tokio::test]
|
||||||
async fn origin_not_allowed() {
|
async fn only_post_and_options_allowed() {
|
||||||
let mut svc = crate::config()
|
|
||||||
.allow_origins(vec!["http://localhost"])
|
|
||||||
.enable(Svc);
|
|
||||||
|
|
||||||
let res = svc.call(request()).await.unwrap();
|
|
||||||
|
|
||||||
assert_eq!(res.status(), StatusCode::FORBIDDEN)
|
|
||||||
}
|
|
||||||
|
|
||||||
#[tokio::test]
|
|
||||||
async fn only_post_allowed() {
|
|
||||||
let mut svc = crate::enable(Svc);
|
let mut svc = crate::enable(Svc);
|
||||||
|
|
||||||
for method in &[
|
for method in &[
|
||||||
@@ -326,7 +314,6 @@ mod tests {
|
|||||||
Method::PUT,
|
Method::PUT,
|
||||||
Method::DELETE,
|
Method::DELETE,
|
||||||
Method::HEAD,
|
Method::HEAD,
|
||||||
Method::OPTIONS,
|
|
||||||
Method::PATCH,
|
Method::PATCH,
|
||||||
] {
|
] {
|
||||||
let mut req = request();
|
let mut req = request();
|
||||||
@@ -361,127 +348,23 @@ mod tests {
|
|||||||
|
|
||||||
mod options {
|
mod options {
|
||||||
use super::*;
|
use super::*;
|
||||||
use crate::cors::{REQUEST_HEADERS, REQUEST_METHOD};
|
|
||||||
use http::HeaderValue;
|
|
||||||
|
|
||||||
const SUCCESS: StatusCode = StatusCode::NO_CONTENT;
|
|
||||||
|
|
||||||
fn request() -> Request<Body> {
|
fn request() -> Request<Body> {
|
||||||
Request::builder()
|
Request::builder()
|
||||||
.method(Method::OPTIONS)
|
.method(Method::OPTIONS)
|
||||||
.header(ORIGIN, "http://example.com")
|
.header(ORIGIN, "http://example.com")
|
||||||
.header(REQUEST_HEADERS, "x-grpc-web")
|
.header(ACCESS_CONTROL_REQUEST_HEADERS, "x-grpc-web")
|
||||||
.header(REQUEST_METHOD, "POST")
|
.header(ACCESS_CONTROL_REQUEST_METHOD, "POST")
|
||||||
.body(Body::empty())
|
.body(Body::empty())
|
||||||
.unwrap()
|
.unwrap()
|
||||||
}
|
}
|
||||||
|
|
||||||
#[tokio::test]
|
|
||||||
async fn origin_not_allowed() {
|
|
||||||
let mut svc = crate::config()
|
|
||||||
.allow_origins(vec!["http://foo.com"])
|
|
||||||
.enable(Svc);
|
|
||||||
|
|
||||||
let res = svc.call(request()).await.unwrap();
|
|
||||||
|
|
||||||
assert_eq!(res.status(), StatusCode::FORBIDDEN);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[tokio::test]
|
|
||||||
async fn missing_request_method() {
|
|
||||||
let mut svc = crate::enable(Svc);
|
|
||||||
|
|
||||||
let mut req = request();
|
|
||||||
req.headers_mut().remove(REQUEST_METHOD);
|
|
||||||
|
|
||||||
let res = svc.call(req).await.unwrap();
|
|
||||||
|
|
||||||
assert_eq!(res.status(), StatusCode::FORBIDDEN);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[tokio::test]
|
|
||||||
async fn only_post_and_options_allowed() {
|
|
||||||
let mut svc = crate::enable(Svc);
|
|
||||||
|
|
||||||
for method in &[
|
|
||||||
Method::GET,
|
|
||||||
Method::PUT,
|
|
||||||
Method::DELETE,
|
|
||||||
Method::HEAD,
|
|
||||||
Method::PATCH,
|
|
||||||
] {
|
|
||||||
let mut req = request();
|
|
||||||
req.headers_mut().insert(
|
|
||||||
REQUEST_METHOD,
|
|
||||||
HeaderValue::from_maybe_shared(method.to_string()).unwrap(),
|
|
||||||
);
|
|
||||||
|
|
||||||
let res = svc.call(req).await.unwrap();
|
|
||||||
|
|
||||||
assert_eq!(
|
|
||||||
res.status(),
|
|
||||||
StatusCode::FORBIDDEN,
|
|
||||||
"{} should not be allowed",
|
|
||||||
method
|
|
||||||
);
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
#[tokio::test]
|
|
||||||
async fn h1_missing_origin_is_err() {
|
|
||||||
let mut svc = crate::enable(Svc);
|
|
||||||
let mut req = request();
|
|
||||||
req.headers_mut().remove(ORIGIN);
|
|
||||||
|
|
||||||
let res = svc.call(req).await.unwrap();
|
|
||||||
|
|
||||||
assert_eq!(res.status(), StatusCode::BAD_REQUEST);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[tokio::test]
|
|
||||||
async fn h2_missing_origin_is_ok() {
|
|
||||||
let mut svc = crate::enable(Svc);
|
|
||||||
|
|
||||||
let mut req = request();
|
|
||||||
*req.version_mut() = Version::HTTP_2;
|
|
||||||
req.headers_mut().remove(ORIGIN);
|
|
||||||
|
|
||||||
let res = svc.call(req).await.unwrap();
|
|
||||||
|
|
||||||
assert_eq!(res.status(), StatusCode::OK);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[tokio::test]
|
|
||||||
async fn h1_missing_x_grpc_web_header_is_err() {
|
|
||||||
let mut svc = crate::enable(Svc);
|
|
||||||
|
|
||||||
let mut req = request();
|
|
||||||
req.headers_mut().remove(REQUEST_HEADERS);
|
|
||||||
|
|
||||||
let res = svc.call(req).await.unwrap();
|
|
||||||
|
|
||||||
assert_eq!(res.status(), StatusCode::BAD_REQUEST);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[tokio::test]
|
|
||||||
async fn h2_missing_x_grpc_web_header_is_ok() {
|
|
||||||
let mut svc = crate::enable(Svc);
|
|
||||||
|
|
||||||
let mut req = request();
|
|
||||||
*req.version_mut() = Version::HTTP_2;
|
|
||||||
req.headers_mut().remove(REQUEST_HEADERS);
|
|
||||||
|
|
||||||
let res = svc.call(req).await.unwrap();
|
|
||||||
|
|
||||||
assert_eq!(res.status(), StatusCode::OK);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[tokio::test]
|
#[tokio::test]
|
||||||
async fn valid_grpc_web_preflight() {
|
async fn valid_grpc_web_preflight() {
|
||||||
let mut svc = crate::enable(Svc);
|
let mut svc = crate::enable(Svc);
|
||||||
let res = svc.call(request()).await.unwrap();
|
let res = svc.call(request()).await.unwrap();
|
||||||
|
|
||||||
assert_eq!(res.status(), SUCCESS);
|
assert_eq!(res.status(), StatusCode::OK);
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -11,6 +11,7 @@ use tonic::{Response, Streaming};
|
|||||||
|
|
||||||
use integration::pb::{test_client::TestClient, test_server::TestServer, Input};
|
use integration::pb::{test_client::TestClient, test_server::TestServer, Input};
|
||||||
use integration::Svc;
|
use integration::Svc;
|
||||||
|
use tonic_web::GrpcWebLayer;
|
||||||
|
|
||||||
#[tokio::test]
|
#[tokio::test]
|
||||||
async fn smoke_unary() {
|
async fn smoke_unary() {
|
||||||
@@ -113,13 +114,10 @@ async fn grpc(accept_h1: bool) -> (impl Future<Output = Result<(), Error>>, Stri
|
|||||||
async fn grpc_web(accept_h1: bool) -> (impl Future<Output = Result<(), Error>>, String) {
|
async fn grpc_web(accept_h1: bool) -> (impl Future<Output = Result<(), Error>>, String) {
|
||||||
let (listener, url) = bind().await;
|
let (listener, url) = bind().await;
|
||||||
|
|
||||||
let svc = tonic_web::config()
|
|
||||||
.allow_origins(vec!["http://foo.com"])
|
|
||||||
.enable(TestServer::new(Svc));
|
|
||||||
|
|
||||||
let fut = Server::builder()
|
let fut = Server::builder()
|
||||||
.accept_http1(accept_h1)
|
.accept_http1(accept_h1)
|
||||||
.add_service(svc)
|
.layer(GrpcWebLayer::new())
|
||||||
|
.add_service(TestServer::new(Svc))
|
||||||
.serve_with_incoming(TcpListenerStream::new(listener));
|
.serve_with_incoming(TcpListenerStream::new(listener));
|
||||||
|
|
||||||
(fut, url)
|
(fut, url)
|
||||||
|
|||||||
@@ -10,10 +10,11 @@ use tonic::transport::Server;
|
|||||||
|
|
||||||
use integration::pb::{test_server::TestServer, Input, Output};
|
use integration::pb::{test_server::TestServer, Input, Output};
|
||||||
use integration::Svc;
|
use integration::Svc;
|
||||||
|
use tonic_web::GrpcWebLayer;
|
||||||
|
|
||||||
#[tokio::test]
|
#[tokio::test]
|
||||||
async fn binary_request() {
|
async fn binary_request() {
|
||||||
let server_url = spawn("http://example.com").await;
|
let server_url = spawn().await;
|
||||||
let client = Client::new();
|
let client = Client::new();
|
||||||
|
|
||||||
let req = build_request(server_url, "grpc-web", "grpc-web");
|
let req = build_request(server_url, "grpc-web", "grpc-web");
|
||||||
@@ -36,7 +37,7 @@ async fn binary_request() {
|
|||||||
|
|
||||||
#[tokio::test]
|
#[tokio::test]
|
||||||
async fn text_request() {
|
async fn text_request() {
|
||||||
let server_url = spawn("http://example.com").await;
|
let server_url = spawn().await;
|
||||||
let client = Client::new();
|
let client = Client::new();
|
||||||
|
|
||||||
let req = build_request(server_url, "grpc-web-text", "grpc-web-text");
|
let req = build_request(server_url, "grpc-web-text", "grpc-web-text");
|
||||||
@@ -57,31 +58,17 @@ async fn text_request() {
|
|||||||
assert_eq!(&trailers[..], b"grpc-status:0\r\n");
|
assert_eq!(&trailers[..], b"grpc-status:0\r\n");
|
||||||
}
|
}
|
||||||
|
|
||||||
#[tokio::test]
|
async fn spawn() -> String {
|
||||||
async fn origin_not_allowed() {
|
|
||||||
let server_url = spawn("http://foo.com").await;
|
|
||||||
let client = Client::new();
|
|
||||||
|
|
||||||
let req = build_request(server_url, "grpc-web-text", "grpc-web-text");
|
|
||||||
let res = client.request(req).await.unwrap();
|
|
||||||
|
|
||||||
assert_eq!(res.status(), StatusCode::FORBIDDEN);
|
|
||||||
}
|
|
||||||
|
|
||||||
async fn spawn(allowed_origin: &str) -> String {
|
|
||||||
let addr = SocketAddr::from(([127, 0, 0, 1], 0));
|
let addr = SocketAddr::from(([127, 0, 0, 1], 0));
|
||||||
let listener = TcpListener::bind(addr).await.expect("listener");
|
let listener = TcpListener::bind(addr).await.expect("listener");
|
||||||
let url = format!("http://{}", listener.local_addr().unwrap());
|
let url = format!("http://{}", listener.local_addr().unwrap());
|
||||||
let listener_stream = TcpListenerStream::new(listener);
|
let listener_stream = TcpListenerStream::new(listener);
|
||||||
|
|
||||||
let svc = tonic_web::config()
|
|
||||||
.allow_origins(vec![allowed_origin])
|
|
||||||
.enable(TestServer::new(Svc));
|
|
||||||
|
|
||||||
let _ = tokio::spawn(async move {
|
let _ = tokio::spawn(async move {
|
||||||
Server::builder()
|
Server::builder()
|
||||||
.accept_http1(true)
|
.accept_http1(true)
|
||||||
.add_service(svc)
|
.layer(GrpcWebLayer::new())
|
||||||
|
.add_service(TestServer::new(Svc))
|
||||||
.serve_with_incoming(listener_stream)
|
.serve_with_incoming(listener_stream)
|
||||||
.await
|
.await
|
||||||
.unwrap()
|
.unwrap()
|
||||||
|
|||||||
Reference in New Issue
Block a user