From afb32c607633d5d8ff3165a82db71ad3f2604a9c Mon Sep 17 00:00:00 2001 From: adelnizamutdinov Date: Wed, 30 Sep 2026 23:50:03 +0300 Subject: [PATCH 1/2] Decode Rust operation requests into typed HTTP response enums --- .gitignore | 1 + README.md | 29 ++++-- internal/rustemit/rustemit.go | 133 ++++++++++++++++++++++++-- moon.yml | 2 + testdata/fixtures/rust-responses.yaml | 28 ++++++ testdata/rust-client/src/lib.rs | 89 ++++++++++++----- 6 files changed, 241 insertions(+), 41 deletions(-) create mode 100644 testdata/fixtures/rust-responses.yaml diff --git a/.gitignore b/.gitignore index 5051869..70eff33 100644 --- a/.gitignore +++ b/.gitignore @@ -16,6 +16,7 @@ dump-*.json /testdata/rust-client/target/ /testdata/rust-client/src/public/ /testdata/rust-client/src/models/ +/testdata/rust-client/src/responses/ /testdata/dart-client/.dart_tool/ /testdata/dart-client/lib/public/ /testdata/dart-client/lib/models/ diff --git a/README.md b/README.md index 08b89e8..8b7700e 100644 --- a/README.md +++ b/README.md @@ -162,6 +162,8 @@ middleware once; every generated operation's `.send()` then propagates the activ trace automatically: ```rust +use api::Response as _; + let http = reqwest_middleware::ClientBuilder::new(reqwest::Client::builder().build()?) .with(reqwest_tracing::TracingMiddleware::default()) .build(); @@ -169,8 +171,11 @@ let http = reqwest_middleware::ClientBuilder::new(reqwest::Client::builder().bui let api = api::Client::new(http, "https://api.example.com".into(), None); // Inside the application's existing tracing span: -let response = api.create_thing(params).send().await?; -let result = api::CreateThingResponse::decode(response).await?; +match api.create_thing(params).send().await? { + api::CreateThingResponse::Status201(thing) => println!("{}", thing.name), + api::CreateThingResponse::Status400(problem) => println!("{}", problem.message), + response => return Err(response.into_error().await.into()), +} ``` For the generated Reqwest 0.12 client, use `reqwest-middleware` 0.4.2 with its @@ -223,11 +228,21 @@ Choose the Reqwest TLS features appropriate to your application. Construct `api::Client::new(http, base_url, bearer_token)` with your configured `reqwest_middleware::ClientWithMiddleware`. For a client without middleware, construct it with `reqwest_middleware::ClientBuilder::new(http).build()`. -Operation methods return `reqwest_middleware::RequestBuilder`, so callers own timeouts, -cancellation, tracing, and bounded body reads. Models and operation-specific -`Response::decode` enums preserve declared HTTP statuses; undeclared statuses -retain their original response. SSE responses remain streaming Reqwest responses, -leaving event framing and cancellation to the caller. +Operation methods return `Request`. Calling `.send().await?` +returns the operation-specific response enum, with typed JSON, raw bytes or an +empty body for every declared status. Match the variants you want to handle; +`Response::into_error()` explicitly turns another variant into a diagnostic error. +Import the generated `Response` trait to use `status()` and `into_error()`. +Undeclared statuses retain the original Reqwest response in `Unexpected`. +SSE responses also remain unread Reqwest responses, leaving event framing and +cancellation to the caller. + +JSON and raw bodies are buffered up to 4 MiB by default; `.body_limit(bytes)` +changes the cap. Invalid JSON and oversized bodies return decode/limit errors. +Use `.send_with(|request| application_send(request))` to retain application +transport policy while decoding the same typed response. Cancellation should +wrap the whole send future so it also interrupts body reads. `.into_builder()` +provides raw HTTP access when needed. See [Rust trace propagation](#rust) to configure automatic propagation once. Rust supports JSON and raw request bodies, optional bodies, scalar and repeated diff --git a/internal/rustemit/rustemit.go b/internal/rustemit/rustemit.go index ebdacf9..19169ca 100644 --- a/internal/rustemit/rustemit.go +++ b/internal/rustemit/rustemit.go @@ -21,6 +21,7 @@ type Options struct{ OutDir string } type emitter struct { schemas map[string]*openapi.Schema definitions map[string]string + serdeModels bool err error } @@ -49,7 +50,11 @@ func Emit(doc *openapi.Document, opts Options, client bool) error { return e.err } var out strings.Builder - out.WriteString("// Code generated by oasmith. DO NOT EDIT.\n#![allow(dead_code, clippy::large_enum_variant, clippy::enum_variant_names)]\nuse serde::{Deserialize, Serialize};\n\n") + out.WriteString("// Code generated by oasmith. DO NOT EDIT.\n#![allow(dead_code, clippy::large_enum_variant, clippy::enum_variant_names)]\n") + if e.serdeModels { + out.WriteString("use serde::{Deserialize, Serialize};\n") + } + out.WriteString("\n") names := make([]string, 0, len(e.definitions)) for name := range e.definitions { names = append(names, name) @@ -158,6 +163,7 @@ func (e *emitter) define(name string, schema *openapi.Schema) { e.definitions[name] = "pub type " + name + " = serde_json::Value;\n" return } + e.serdeModels = true out.WriteString("#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)]\npub enum " + name + " {\n") used := map[string]bool{} for index, value := range values { @@ -175,6 +181,7 @@ func (e *emitter) define(name string, schema *openapi.Schema) { case schema != nil && len(schema.OneOf) > 0 && schema.Discriminator != nil: e.taggedUnion(&out, name, schema, nil) case schema != nil && len(schema.OneOf) > 0: + e.serdeModels = true out.WriteString("#[derive(Clone, Debug, Serialize, Deserialize)]\n#[serde(untagged)]\npub enum " + name + " {\n") for index, item := range schema.OneOf { typ := e.rustType(item, fmt.Sprintf("%sVariant%d", name, index)) @@ -233,6 +240,7 @@ func (e *emitter) taggedUnion(out *strings.Builder, name string, schema *openapi allowed, ok := item.schema.AdditionalProperties.(bool) strict = strict && ok && !allowed } + e.serdeModels = true fmt.Fprintf(out, "#[derive(Clone, Debug, Serialize, Deserialize)]\n#[serde(tag = %s", strconv.Quote(tag)) if strict { out.WriteString(", deny_unknown_fields") @@ -303,6 +311,7 @@ func (e *emitter) fields(out *strings.Builder, name string, schema *openapi.Sche } func (e *emitter) object(out *strings.Builder, name string, schema *openapi.Schema) { + e.serdeModels = true out.WriteString("#[derive(Clone, Debug, Serialize, Deserialize)]\n") if allowed, ok := schema.AdditionalProperties.(bool); ok && !allowed { out.WriteString("#[serde(deny_unknown_fields)]\n") @@ -334,7 +343,7 @@ func (e *emitter) operation(op clientgen.Operation) string { if hasParams { fmt.Fprintf(&body, ", params: %sParams", name) } - body.WriteString(") -> reqwest_middleware::RequestBuilder {\n") + fmt.Fprintf(&body, ") -> Request<%sResponse> {\n", name) fmt.Fprintf(&body, " let path = %s.to_owned();\n", strconv.Quote(op.Route.Path)) for _, param := range op.Route.Operation.Parameters { typ := e.rustType(param.Schema, name+typeName(param.Name)) @@ -374,9 +383,8 @@ func (e *emitter) operation(op clientgen.Operation) string { e.definitions[name+"Params"] = params.String() } e.response(op, name) - source := body.String() - last := strings.LastIndex(source, " let request = ") - return source[:last] + " " + strings.TrimSuffix(source[last+len(" let request = "):], ";\n") + "\n }\n" + body.WriteString(" Request::new(request)\n }\n") + return body.String() } func (e *emitter) parameters(body *strings.Builder, op clientgen.Operation) { @@ -423,24 +431,131 @@ func (e *emitter) response(op clientgen.Operation, name string) { fmt.Fprintf(&response, " Status%d(%s),\n", res.Status, typ) } response.WriteString(" Unexpected(reqwest::Response),\n}\n") - fmt.Fprintf(&response, "impl %sResponse {\n pub async fn decode(response: reqwest::Response) -> Result {\n Ok(match response.status().as_u16() {\n", name) + fmt.Fprintf(&response, "impl Response for %sResponse {\n async fn decode(response: reqwest::Response, _limit: usize) -> Result {\n Ok(match response.status().as_u16() {\n", name) for _, res := range op.Responses { - value := "response.json().await?" + value := "serde_json::from_slice(&read_body(response, _limit).await?).map_err(Error::Decode)?" switch { case res.SSE: value = "response" case !res.HasBody(): value = "()" case !res.JSON(): - value = "response.bytes().await?.to_vec()" + value = "read_body(response, _limit).await?" } fmt.Fprintf(&response, " %d => Self::Status%d(%s),\n", res.Status, res.Status, value) } - response.WriteString(" _ => Self::Unexpected(response),\n })\n }\n}\n") + response.WriteString(" _ => Self::Unexpected(response),\n })\n }\n fn status(&self) -> u16 {\n match self {\n") + for _, res := range op.Responses { + fmt.Fprintf(&response, " Self::Status%d(_) => %d,\n", res.Status, res.Status) + } + response.WriteString(" Self::Unexpected(response) => response.status().as_u16(),\n }\n }\n") + if slices.ContainsFunc(op.Responses, func(res clientgen.Response) bool { return res.SSE }) { + response.WriteString(" fn into_stream(self) -> Result {\n match self {\n") + for _, res := range op.Responses { + if res.SSE { + fmt.Fprintf(&response, " Self::Status%d(response) => Ok(response),\n", res.Status) + } + } + response.WriteString(" response => Err(response),\n }\n }\n") + } + response.WriteString(" async fn into_error(self) -> Error {\n let status = self.status();\n let body = match self {\n") + for _, res := range op.Responses { + pattern, value := "body", "serde_json::to_vec(&body).map_err(Error::Decode)" + switch { + case res.SSE: + value = "read_body(body, 64 << 10).await" + case !res.HasBody(): + pattern, value = "_", "Ok(Vec::new())" + case !res.JSON(): + value = "Ok(body)" + } + fmt.Fprintf(&response, " Self::Status%d(%s) => %s,\n", res.Status, pattern, value) + } + response.WriteString(" Self::Unexpected(response) => read_body(response, 64 << 10).await,\n };\n match body { Ok(body) => Error::Http { status, body }, Err(error) => error }\n }\n}\n") e.definitions[name+"Response"] = response.String() } const clientPrelude = ` +/// A declared response is decoded exactly once; unknown and streaming bodies stay unread. +pub trait Response: Sized + Send { + fn decode(response: reqwest::Response, limit: usize) -> impl std::future::Future> + Send; + fn status(&self) -> u16; + fn into_stream(self) -> Result { Err(self) } + /// Convert a response to an error only after the caller chooses not to handle it. + fn into_error(self) -> impl std::future::Future + Send; +} + +#[derive(Debug)] +pub enum Error { + Transport(reqwest_middleware::Error), + Body(reqwest::Error), + Decode(serde_json::Error), + BodyTooLarge { limit: usize }, + Http { status: u16, body: Vec }, +} +impl std::fmt::Display for Error { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + match self { + Self::Transport(error) => write!(f, "{error}"), + Self::Body(error) => write!(f, "{error}"), + Self::Decode(error) => write!(f, "decode API response: {error}"), + Self::BodyTooLarge { limit } => write!(f, "HTTP response exceeds {limit} bytes"), + Self::Http { status, body } => write!(f, "HTTP {status}: {}", String::from_utf8_lossy(body)), + } + } +} +impl std::error::Error for Error { + fn source(&self) -> Option<&(dyn std::error::Error + 'static)> { + match self { + Self::Transport(error) => Some(error), + Self::Body(error) => Some(error), + Self::Decode(error) => Some(error), + _ => None, + } + } +} +impl From for Error { + fn from(error: reqwest_middleware::Error) -> Self { Self::Transport(error.without_url()) } +} +async fn read_body(mut response: reqwest::Response, limit: usize) -> Result, Error> { + let mut body = Vec::new(); + while let Some(chunk) = response.chunk().await.map_err(|error| Error::Body(error.without_url()))? { + if chunk.len() > limit.saturating_sub(body.len()) { return Err(Error::BodyTooLarge { limit }); } + body.extend_from_slice(&chunk); + } + Ok(body) +} + +/// The response type is fixed by the OpenAPI operation, including every declared status. +#[derive(Debug)] +pub struct Request { + builder: reqwest_middleware::RequestBuilder, + limit: usize, + response: std::marker::PhantomData R>, +} +impl Request { + fn new(builder: reqwest_middleware::RequestBuilder) -> Self { + Self { builder, limit: 4 << 20, response: std::marker::PhantomData } + } + /// Cap buffered JSON and raw bodies. Streaming and unknown responses stay unread. + pub fn body_limit(mut self, limit: usize) -> Self { self.limit = limit; self } + /// Access the underlying builder for transport configuration or raw HTTP handling. + pub fn into_builder(self) -> reqwest_middleware::RequestBuilder { self.builder } + pub fn build(self) -> reqwest::Result { self.builder.build() } + pub async fn send(self) -> Result { + R::decode(self.builder.send().await.map_err(Error::from)?, self.limit).await + } + /// Keep application transport policy while using the operation's generated decoder. + pub async fn send_with(self, send: F) -> Result + where + E: From, + F: FnOnce(reqwest_middleware::RequestBuilder) -> Fut, + Fut: std::future::Future>, + { + R::decode(send(self.builder).await?, self.limit).await.map_err(E::from) + } +} + #[derive(Clone, Debug)] pub struct Client { http: reqwest_middleware::ClientWithMiddleware, diff --git a/moon.yml b/moon.yml index 68e487d..cdbc937 100644 --- a/moon.yml +++ b/moon.yml @@ -68,9 +68,11 @@ tasks: set -euo pipefail go run ./cmd/oasmith --openapi testdata/fixtures/public-client.yaml --mode client --lang rust --out testdata/rust-client/src/public go run ./cmd/oasmith --openapi testdata/fixtures/rust-models.yaml --mode types --lang rust --out testdata/rust-client/src/models + go run ./cmd/oasmith --openapi testdata/fixtures/rust-responses.yaml --mode client --lang rust --out testdata/rust-client/src/responses outputs: - testdata/rust-client/src/public/mod.rs - testdata/rust-client/src/models/mod.rs + - testdata/rust-client/src/responses/mod.rs options: cache: true test-rust: diff --git a/testdata/fixtures/rust-responses.yaml b/testdata/fixtures/rust-responses.yaml new file mode 100644 index 0000000..8943c03 --- /dev/null +++ b/testdata/fixtures/rust-responses.yaml @@ -0,0 +1,28 @@ +openapi: 3.2.0 +info: + title: Rust status response fixture + version: "1" +paths: + /payload: + get: + operationId: getPayload + responses: + "200": + description: JSON text + content: + application/json: + schema: + type: string + "202": + description: JSON number + content: + application/json: + schema: + type: integer + "206": + description: Raw media + content: + application/octet-stream: + schema: + type: string + format: binary diff --git a/testdata/rust-client/src/lib.rs b/testdata/rust-client/src/lib.rs index 66b5c48..d6fd385 100644 --- a/testdata/rust-client/src/lib.rs +++ b/testdata/rust-client/src/lib.rs @@ -1,9 +1,10 @@ mod public; mod models; +mod responses; #[cfg(test)] mod tests { - use super::{public::*, models}; + use super::{public::*, models, responses}; fn client() -> Client { let http = reqwest_middleware::ClientBuilder::new(reqwest::Client::new()).build(); Client::new(http, "http://localhost:1234/".into(), Some("test-key".into())) @@ -28,8 +29,7 @@ mod tests { } #[test] fn optional_bodies_are_absent_and_raw_bytes_are_unchanged() { - for request in [client().patch_thing(PatchThingParams {body: None}), client().upload_optional_media(UploadOptionalMediaParams {body: None})] { - let request = request.build().unwrap(); + for request in [client().patch_thing(PatchThingParams {body: None}).build().unwrap(), client().upload_optional_media(UploadOptionalMediaParams {body: None}).build().unwrap()] { assert!(request.body().is_none()); assert!(!request.headers().contains_key("content-type")); } @@ -38,22 +38,64 @@ mod tests { assert_eq!(request.body().unwrap().as_bytes().unwrap(), &[0, 255, 10]); } #[tokio::test] - async fn decodes_statuses_and_leaves_event_stream_unconsumed() { - let response = reqwest::Response::from(http_response(201, r#"{"id":"a","name":"Podcast"}"#)); - let CreateThingResponse::Status201(thing) = CreateThingResponse::decode(response).await.unwrap() else { panic!("wrong success variant") }; - assert_eq!(thing.name, "Podcast"); - let response = reqwest::Response::from(http_response(400, r#"{"message":"invalid"}"#)); - assert!(matches!(CreateThingResponse::decode(response).await.unwrap(), CreateThingResponse::Status400(_))); - let response = reqwest::Response::from(http_response(204, "")); - assert!(matches!(PatchThingResponse::decode(response).await.unwrap(), PatchThingResponse::Status204(()))); - let response = reqwest::Response::from(http_response(502, "unavailable")); - assert!(matches!(CreateThingResponse::decode(response).await.unwrap(), CreateThingResponse::Unexpected(_))); - let request = client().watch_events().build().unwrap(); - assert_eq!(request.headers()["accept"], "text/event-stream"); - let event = "event: progress\ndata: {}\n\n"; - let response = reqwest::Response::from(http_response(200, event)); - let WatchEventsResponse::Status200(stream) = WatchEventsResponse::decode(response).await.unwrap() else { panic!("wrong event variant") }; - assert_eq!(stream.text().await.unwrap(), event); + async fn generated_operations_decode_http_statuses_and_bodies() { + use wiremock::{Mock, MockServer, ResponseTemplate, matchers::{method, path}}; + let work = async { + let server = MockServer::start().await; + let http = reqwest_middleware::ClientBuilder::new(reqwest::Client::builder().no_proxy().build().unwrap()).build(); + let api = Client::new(http, server.uri(), None); + for (status, body) in [(201, r#"{"id":"a","name":"Podcast"}"#), (400, r#"{"message":"invalid"}"#), (401, ""), (403, ""), (502, "unavailable")] { + Mock::given(method("POST")).and(path("/things/a")) + .respond_with(ResponseTemplate::new(status).set_body_string(body)).expect(1).mount(&server).await; + let result = api.create_thing(CreateThingParams { + thing_id: "a".into(), tag: None, notify: false, label: None, + x_request_id: "request".into(), body: CreateThing { name: "Podcast".into() }, + }).send().await.unwrap(); + assert_eq!(result.status(), status); + match (status, result) { + (201, CreateThingResponse::Status201(thing)) => assert_eq!(thing.name, "Podcast"), + (400, CreateThingResponse::Status400(problem)) => assert_eq!(problem.message, "invalid"), + (401, CreateThingResponse::Status401(())) | (403, CreateThingResponse::Status403(())) => {}, + (502, CreateThingResponse::Unexpected(response)) => assert_eq!(response.text().await.unwrap(), "unavailable"), + _ => panic!("wrong response variant for {status}"), + } + server.verify().await; + server.reset().await; + } + Mock::given(path("/optional-json")).respond_with(ResponseTemplate::new(204)).expect(1).mount(&server).await; + assert!(matches!(api.patch_thing(PatchThingParams { body: None }).send().await.unwrap(), PatchThingResponse::Status204(()))); + for (body, oversized) in [("invalid JSON", false), (r#"{"id":"a","name":"Podcast"}"#, true)] { + Mock::given(path("/things/a")).respond_with(ResponseTemplate::new(201).set_body_string(body)).expect(1).mount(&server).await; + let result = api.create_thing(CreateThingParams { + thing_id: "a".into(), tag: None, notify: false, label: None, + x_request_id: "request".into(), body: CreateThing { name: "Podcast".into() }, + }).body_limit(if oversized { 8 } else { 1024 }).send().await; + if oversized { assert!(matches!(result, Err(Error::BodyTooLarge { .. }))); } + else { assert!(matches!(result, Err(Error::Decode(_)))); } + server.verify().await; + server.reset().await; + } + let payloads = responses::Client::new(reqwest_middleware::ClientBuilder::new(reqwest::Client::builder().no_proxy().build().unwrap()).build(), server.uri(), None); + for (status, body) in [(200, r#""Podcast""#), (202, "42"), (206, "raw media")] { + Mock::given(path("/payload")).respond_with(ResponseTemplate::new(status).set_body_string(body)).expect(1).mount(&server).await; + let result = payloads.get_payload().send().await.unwrap(); + assert_eq!(responses::Response::status(&result), status); + match result { + responses::GetPayloadResponse::Status200(text) => assert_eq!(text, "Podcast"), + responses::GetPayloadResponse::Status202(number) => assert_eq!(number, 42), + responses::GetPayloadResponse::Status206(bytes) => assert_eq!(bytes, b"raw media"), + _ => panic!("wrong response variant for {status}"), + } + server.verify().await; + server.reset().await; + } + let event = "event: progress\ndata: {}\n\n"; + Mock::given(path("/events")).respond_with(ResponseTemplate::new(200).set_body_string(event)).expect(1).mount(&server).await; + let WatchEventsResponse::Status200(stream) = api.watch_events().send().await.unwrap() else { panic!("wrong event variant") }; + assert_eq!(stream.text().await.unwrap(), event); + server.verify().await; + }; + tokio::time::timeout(std::time::Duration::from_secs(5), work).await.expect("typed HTTP contract did not complete"); } #[test] fn discriminated_oneof_is_a_named_tagged_enum() { @@ -78,9 +120,6 @@ mod tests { assert_eq!(wire, serde_json::json!({"type":"ImplicitPayload","value":"data"})); assert!(matches!(serde_json::from_value::(wire).unwrap(), models::ImplicitEvent::ImplicitPayload { .. })); } - fn http_response(status: u16, body: &str) -> http::Response { - http::Response::builder().status(status).body(body.into()).unwrap() - } #[test] fn models_preserve_wire_names_nullability_and_discriminators() { let value: models::Record = serde_json::from_value(serde_json::json!({"type":"podcast", "nullable":null, "self":"me"})).unwrap(); @@ -145,9 +184,9 @@ mod tests { ))); } let parents = [first.context().span().span_context().clone(), second.context().span().span_context().clone()]; - let (post, events) = tokio::join!(post.send().instrument(first), events.send().instrument(second)); - assert!(matches!(CreateThingResponse::decode(post.unwrap()).await.unwrap(), CreateThingResponse::Status400(_))); - let WatchEventsResponse::Status200(stream) = WatchEventsResponse::decode(events.unwrap()).await.unwrap() else { panic!("wrong event variant") }; + let (post, events) = tokio::join!(post.send_with(|request| async { request.send().await.map_err(Error::from) }).instrument(first), events.send().instrument(second)); + assert!(matches!(post.unwrap(), CreateThingResponse::Status400(_))); + let WatchEventsResponse::Status200(stream) = events.unwrap() else { panic!("wrong event variant") }; assert_eq!(stream.text().await.unwrap(), event); server.verify().await; From 5d0e82bb731ec930c8b6899e7151c1da1f6dfba6 Mon Sep 17 00:00:00 2001 From: adelnizamutdinov Date: Thu, 1 Oct 2026 05:47:03 +0300 Subject: [PATCH 2/2] Generate directly awaitable Rust API operations --- README.md | 16 ++--- internal/rustemit/rustemit.go | 56 ++++++++--------- testdata/rust-client/src/lib.rs | 105 ++++++++++++++++++++------------ 3 files changed, 99 insertions(+), 78 deletions(-) diff --git a/README.md b/README.md index 8b7700e..0be59d9 100644 --- a/README.md +++ b/README.md @@ -158,7 +158,7 @@ by the application, commonly `traceparent`, `tracestate`, and `baggage`. The insertion point is the **HTTP client passed to `api::Client::new`**. Configure [`reqwest-tracing`](https://docs.rs/reqwest-tracing/0.5.8/reqwest_tracing/) -middleware once; every generated operation's `.send()` then propagates the active +middleware once; every awaited generated operation propagates the active trace automatically: ```rust @@ -171,7 +171,7 @@ let http = reqwest_middleware::ClientBuilder::new(reqwest::Client::builder().bui let api = api::Client::new(http, "https://api.example.com".into(), None); // Inside the application's existing tracing span: -match api.create_thing(params).send().await? { +match api.create_thing(params).await? { api::CreateThingResponse::Status201(thing) => println!("{}", thing.name), api::CreateThingResponse::Status400(problem) => println!("{}", problem.message), response => return Err(response.into_error().await.into()), @@ -228,8 +228,8 @@ Choose the Reqwest TLS features appropriate to your application. Construct `api::Client::new(http, base_url, bearer_token)` with your configured `reqwest_middleware::ClientWithMiddleware`. For a client without middleware, construct it with `reqwest_middleware::ClientBuilder::new(http).build()`. -Operation methods return `Request`. Calling `.send().await?` -returns the operation-specific response enum, with typed JSON, raw bytes or an +Await an operation directly: `api.operation(params).await?` returns its +operation-specific response enum, with typed JSON, raw bytes or an empty body for every declared status. Match the variants you want to handle; `Response::into_error()` explicitly turns another variant into a diagnostic error. Import the generated `Response` trait to use `status()` and `into_error()`. @@ -239,10 +239,10 @@ cancellation to the caller. JSON and raw bodies are buffered up to 4 MiB by default; `.body_limit(bytes)` changes the cap. Invalid JSON and oversized bodies return decode/limit errors. -Use `.send_with(|request| application_send(request))` to retain application -transport policy while decoding the same typed response. Cancellation should -wrap the whole send future so it also interrupts body reads. `.into_builder()` -provides raw HTTP access when needed. +Configure application transport policy once with `Client::with_transport`. +Implement the generated `Transport` trait to apply cancellation or diagnostics +across sending and typed response decoding. Operations still take only their +parameters and return their response enum directly. See [Rust trace propagation](#rust) to configure automatic propagation once. Rust supports JSON and raw request bodies, optional bodies, scalar and repeated diff --git a/internal/rustemit/rustemit.go b/internal/rustemit/rustemit.go index 19169ca..ceb6361 100644 --- a/internal/rustemit/rustemit.go +++ b/internal/rustemit/rustemit.go @@ -338,12 +338,12 @@ func (e *emitter) operation(op clientgen.Operation) string { name := typeName(op.Route.Operation.OperationID) var params, body strings.Builder params.WriteString("#[derive(Clone, Debug)]\npub struct " + name + "Params {\n") - fmt.Fprintf(&body, " pub fn %s(&self", snake(op.Route.Operation.OperationID)) + fmt.Fprintf(&body, " pub async fn %s(&self", snake(op.Route.Operation.OperationID)) hasParams := len(op.Route.Operation.Parameters) > 0 || op.RequestBody.JSON != nil || op.RequestBody.Raw != nil if hasParams { fmt.Fprintf(&body, ", params: %sParams", name) } - fmt.Fprintf(&body, ") -> Request<%sResponse> {\n", name) + fmt.Fprintf(&body, ") -> Result<%sResponse, T::Error> {\n", name) fmt.Fprintf(&body, " let path = %s.to_owned();\n", strconv.Quote(op.Route.Path)) for _, param := range op.Route.Operation.Parameters { typ := e.rustType(param.Schema, name+typeName(param.Name)) @@ -383,7 +383,7 @@ func (e *emitter) operation(op clientgen.Operation) string { e.definitions[name+"Params"] = params.String() } e.response(op, name) - body.WriteString(" Request::new(request)\n }\n") + body.WriteString(" self.transport.execute(request, self.body_limit).await\n }\n") return body.String() } @@ -526,41 +526,26 @@ async fn read_body(mut response: reqwest::Response, limit: usize) -> Result { - builder: reqwest_middleware::RequestBuilder, - limit: usize, - response: std::marker::PhantomData R>, +/// Configure application transport policy once, including sending and response reads. +pub trait Transport: Sync { + type Error: From + Send; + fn execute(&self, request: reqwest_middleware::RequestBuilder, limit: usize) + -> impl std::future::Future> + Send; } -impl Request { - fn new(builder: reqwest_middleware::RequestBuilder) -> Self { - Self { builder, limit: 4 << 20, response: std::marker::PhantomData } - } - /// Cap buffered JSON and raw bodies. Streaming and unknown responses stay unread. - pub fn body_limit(mut self, limit: usize) -> Self { self.limit = limit; self } - /// Access the underlying builder for transport configuration or raw HTTP handling. - pub fn into_builder(self) -> reqwest_middleware::RequestBuilder { self.builder } - pub fn build(self) -> reqwest::Result { self.builder.build() } - pub async fn send(self) -> Result { - R::decode(self.builder.send().await.map_err(Error::from)?, self.limit).await - } - /// Keep application transport policy while using the operation's generated decoder. - pub async fn send_with(self, send: F) -> Result - where - E: From, - F: FnOnce(reqwest_middleware::RequestBuilder) -> Fut, - Fut: std::future::Future>, - { - R::decode(send(self.builder).await?, self.limit).await.map_err(E::from) +impl Transport for () { + type Error = Error; + async fn execute(&self, request: reqwest_middleware::RequestBuilder, limit: usize) -> Result { + R::decode(request.send().await.map_err(Error::from)?, limit).await } } #[derive(Clone, Debug)] -pub struct Client { +pub struct Client { http: reqwest_middleware::ClientWithMiddleware, base_url: String, bearer_token: Option, + body_limit: usize, + transport: T, } fn encode_path(value: &str) -> String { value.bytes().map(|byte| { @@ -570,8 +555,15 @@ fn encode_path(value: &str) -> String { }).collect() } impl Client { - /// Every operation uses this client's middleware when its request is sent. + /// Operations send and decode directly through this client's middleware. pub fn new(http: reqwest_middleware::ClientWithMiddleware, base_url: String, bearer_token: Option) -> Self { - Self { http, base_url: base_url.trim_end_matches('/').to_owned(), bearer_token } + Self { http, base_url: base_url.trim_end_matches('/').to_owned(), bearer_token, body_limit: 4 << 20, transport: () } + } +} +impl Client { + /// Cap buffered JSON and raw bodies for every operation. Streaming and unknown responses stay unread. + pub fn body_limit(mut self, limit: usize) -> Self { self.body_limit = limit; self } + pub fn with_transport(self, transport: U) -> Client { + Client { http: self.http, base_url: self.base_url, bearer_token: self.bearer_token, body_limit: self.body_limit, transport } } ` diff --git a/testdata/rust-client/src/lib.rs b/testdata/rust-client/src/lib.rs index d6fd385..1b26773 100644 --- a/testdata/rust-client/src/lib.rs +++ b/testdata/rust-client/src/lib.rs @@ -5,37 +5,64 @@ mod responses; #[cfg(test)] mod tests { use super::{public::*, models, responses}; - fn client() -> Client { - let http = reqwest_middleware::ClientBuilder::new(reqwest::Client::new()).build(); - Client::new(http, "http://localhost:1234/".into(), Some("test-key".into())) - } - #[test] - fn encodes_path_repeated_query_headers_and_json() { - let request = client().create_thing(CreateThingParams { - thing_id: "a/b ?#%".into(), tag: Some("x & y".into()), notify: false, - label: Some(vec!["one".into(), "two".into()]), x_request_id: "req-1".into(), - body: CreateThing { name: "Podcast".into() }, - }).build().unwrap(); - assert_eq!(request.url().path(), "/things/a%2Fb%20%3F%23%25"); - assert_eq!(request.url().query_pairs().collect::>(), vec![ - ("tag".into(), "x & y".into()), ("notify".into(), "false".into()), - ("label".into(), "one".into()), ("label".into(), "two".into()), - ]); - assert_eq!(request.headers()["authorization"], "Bearer test-key"); - assert_eq!(request.headers()["x-request-id"], "req-1"); - assert_eq!(request.headers()["content-type"], "application/json"); - let body: serde_json::Value = serde_json::from_slice(request.body().unwrap().as_bytes().unwrap()).unwrap(); - assert_eq!(body, serde_json::json!({"name":"Podcast"})); + #[tokio::test] + async fn operations_send_encoded_paths_queries_headers_and_json() { + use wiremock::{Mock, MockServer, ResponseTemplate, matchers::method}; + let work = async { + let server = MockServer::start().await; + Mock::given(method("POST")) + .respond_with(ResponseTemplate::new(201).set_body_json(serde_json::json!({"id":"a", "name":"Podcast"}))) + .expect(1).mount(&server).await; + let http = reqwest_middleware::ClientBuilder::new(reqwest::Client::builder().no_proxy().build().unwrap()).build(); + let api = Client::new(http, server.uri(), Some("test-key".into())); + let result = api.create_thing(CreateThingParams { + thing_id: "a/b ?#%".into(), tag: Some("x & y".into()), notify: false, + label: Some(vec!["one".into(), "two".into()]), x_request_id: "req-1".into(), + body: CreateThing { name: "Podcast".into() }, + }).await.unwrap(); + assert!(matches!(result, CreateThingResponse::Status201(_))); + server.verify().await; + let requests = server.received_requests().await.unwrap(); + let request = &requests[0]; + assert_eq!(request.url.path(), "/things/a%2Fb%20%3F%23%25"); + assert_eq!(request.url.query_pairs().collect::>(), vec![ + ("tag".into(), "x & y".into()), ("notify".into(), "false".into()), + ("label".into(), "one".into()), ("label".into(), "two".into()), + ]); + assert_eq!(request.headers["authorization"], "Bearer test-key"); + assert_eq!(request.headers["x-request-id"], "req-1"); + assert_eq!(request.headers["content-type"], "application/json"); + assert_eq!(serde_json::from_slice::(&request.body).unwrap(), serde_json::json!({"name":"Podcast"})); + }; + tokio::time::timeout(std::time::Duration::from_secs(5), work).await.expect("HTTP input contract did not complete"); } - #[test] - fn optional_bodies_are_absent_and_raw_bytes_are_unchanged() { - for request in [client().patch_thing(PatchThingParams {body: None}).build().unwrap(), client().upload_optional_media(UploadOptionalMediaParams {body: None}).build().unwrap()] { - assert!(request.body().is_none()); - assert!(!request.headers().contains_key("content-type")); - } - let request = client().upload_media(UploadMediaParams {owner: "owner".into(), upload_type: "media".into(), body: vec![0, 255, 10]}).build().unwrap(); - assert_eq!(request.headers()["content-type"], "application/octet-stream"); - assert_eq!(request.body().unwrap().as_bytes().unwrap(), &[0, 255, 10]); + #[tokio::test] + async fn operations_send_optional_bodies_and_raw_bytes() { + use wiremock::{Mock, MockServer, ResponseTemplate, matchers::path}; + let work = async { + let server = MockServer::start().await; + for route in ["/optional-json", "/optional-raw"] { + Mock::given(path(route)).respond_with(ResponseTemplate::new(204)).expect(1).mount(&server).await; + } + Mock::given(path("/uploads/owner")) + .respond_with(ResponseTemplate::new(201).set_body_json(serde_json::json!({"id":"a", "name":"Podcast"}))) + .expect(1).mount(&server).await; + let http = reqwest_middleware::ClientBuilder::new(reqwest::Client::builder().no_proxy().build().unwrap()).build(); + let api = Client::new(http, server.uri(), None); + assert!(matches!(api.patch_thing(PatchThingParams {body: None}).await.unwrap(), PatchThingResponse::Status204(()))); + assert!(matches!(api.upload_optional_media(UploadOptionalMediaParams {body: None}).await.unwrap(), UploadOptionalMediaResponse::Status204(()))); + assert!(matches!(api.upload_media(UploadMediaParams {owner: "owner".into(), upload_type: "media".into(), body: vec![0, 255, 10]}).await.unwrap(), UploadMediaResponse::Status201(_))); + server.verify().await; + let requests = server.received_requests().await.unwrap(); + for request in requests.iter().filter(|r| r.url.path().starts_with("/optional-")) { + assert!(request.body.is_empty()); + assert!(!request.headers.contains_key("content-type")); + } + let request = requests.iter().find(|r| r.url.path() == "/uploads/owner").unwrap(); + assert_eq!(request.headers["content-type"], "application/octet-stream"); + assert_eq!(request.body, &[0, 255, 10]); + }; + tokio::time::timeout(std::time::Duration::from_secs(5), work).await.expect("HTTP body contract did not complete"); } #[tokio::test] async fn generated_operations_decode_http_statuses_and_bodies() { @@ -50,7 +77,7 @@ mod tests { let result = api.create_thing(CreateThingParams { thing_id: "a".into(), tag: None, notify: false, label: None, x_request_id: "request".into(), body: CreateThing { name: "Podcast".into() }, - }).send().await.unwrap(); + }).await.unwrap(); assert_eq!(result.status(), status); match (status, result) { (201, CreateThingResponse::Status201(thing)) => assert_eq!(thing.name, "Podcast"), @@ -63,13 +90,14 @@ mod tests { server.reset().await; } Mock::given(path("/optional-json")).respond_with(ResponseTemplate::new(204)).expect(1).mount(&server).await; - assert!(matches!(api.patch_thing(PatchThingParams { body: None }).send().await.unwrap(), PatchThingResponse::Status204(()))); + assert!(matches!(api.patch_thing(PatchThingParams { body: None }).await.unwrap(), PatchThingResponse::Status204(()))); for (body, oversized) in [("invalid JSON", false), (r#"{"id":"a","name":"Podcast"}"#, true)] { Mock::given(path("/things/a")).respond_with(ResponseTemplate::new(201).set_body_string(body)).expect(1).mount(&server).await; - let result = api.create_thing(CreateThingParams { + let bounded = api.clone().body_limit(if oversized { 8 } else { 1024 }); + let result = bounded.create_thing(CreateThingParams { thing_id: "a".into(), tag: None, notify: false, label: None, x_request_id: "request".into(), body: CreateThing { name: "Podcast".into() }, - }).body_limit(if oversized { 8 } else { 1024 }).send().await; + }).await; if oversized { assert!(matches!(result, Err(Error::BodyTooLarge { .. }))); } else { assert!(matches!(result, Err(Error::Decode(_)))); } server.verify().await; @@ -78,7 +106,7 @@ mod tests { let payloads = responses::Client::new(reqwest_middleware::ClientBuilder::new(reqwest::Client::builder().no_proxy().build().unwrap()).build(), server.uri(), None); for (status, body) in [(200, r#""Podcast""#), (202, "42"), (206, "raw media")] { Mock::given(path("/payload")).respond_with(ResponseTemplate::new(status).set_body_string(body)).expect(1).mount(&server).await; - let result = payloads.get_payload().send().await.unwrap(); + let result = payloads.get_payload().await.unwrap(); assert_eq!(responses::Response::status(&result), status); match result { responses::GetPayloadResponse::Status200(text) => assert_eq!(text, "Podcast"), @@ -91,7 +119,7 @@ mod tests { } let event = "event: progress\ndata: {}\n\n"; Mock::given(path("/events")).respond_with(ResponseTemplate::new(200).set_body_string(event)).expect(1).mount(&server).await; - let WatchEventsResponse::Status200(stream) = api.watch_events().send().await.unwrap() else { panic!("wrong event variant") }; + let WatchEventsResponse::Status200(stream) = api.watch_events().await.unwrap() else { panic!("wrong event variant") }; assert_eq!(stream.text().await.unwrap(), event); server.verify().await; }; @@ -174,7 +202,8 @@ mod tests { thing_id: "a/b".into(), tag: None, notify: false, label: None, x_request_id: "request-one".into(), body: CreateThing { name: "Podcast".into() }, }); - let events = api.clone().watch_events(); + let cloned = api.clone(); + let events = cloned.watch_events(); let first = tracing::info_span!("first operation"); let second = tracing::info_span!("second operation"); for (span, trace_id, span_id) in [(&first, 1u128, 2u64), (&second, 3u128, 4u64)] { @@ -184,7 +213,7 @@ mod tests { ))); } let parents = [first.context().span().span_context().clone(), second.context().span().span_context().clone()]; - let (post, events) = tokio::join!(post.send_with(|request| async { request.send().await.map_err(Error::from) }).instrument(first), events.send().instrument(second)); + let (post, events) = tokio::join!(post.instrument(first), events.instrument(second)); assert!(matches!(post.unwrap(), CreateThingResponse::Status400(_))); let WatchEventsResponse::Status200(stream) = events.unwrap() else { panic!("wrong event variant") }; assert_eq!(stream.text().await.unwrap(), event);