Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
1 change: 1 addition & 0 deletions .gitignore
Original file line number Diff line number Diff line change
Expand Up @@ -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/
31 changes: 23 additions & 8 deletions README.md
Original file line number Diff line number Diff line change
Expand Up @@ -158,19 +158,24 @@ 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
use api::Response as _;

let http = reqwest_middleware::ClientBuilder::new(reqwest::Client::builder().build()?)
.with(reqwest_tracing::TracingMiddleware::default())
.build();

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).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
Expand Down Expand Up @@ -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.
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()`.
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.
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
Expand Down
133 changes: 120 additions & 13 deletions internal/rustemit/rustemit.go
Original file line number Diff line number Diff line change
Expand Up @@ -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
}

Expand Down Expand Up @@ -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)
Expand Down Expand Up @@ -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 {
Expand All @@ -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))
Expand Down Expand Up @@ -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")
Expand Down Expand Up @@ -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")
Expand All @@ -329,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)
}
body.WriteString(") -> reqwest_middleware::RequestBuilder {\n")
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))
Expand Down Expand Up @@ -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(" self.transport.execute(request, self.body_limit).await\n }\n")
return body.String()
}

func (e *emitter) parameters(body *strings.Builder, op clientgen.Operation) {
Expand Down Expand Up @@ -423,29 +431,121 @@ 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<Self, reqwest::Error> {\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<Self, Error> {\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<reqwest::Response, Self> {\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<Output = Result<Self, Error>> + Send;
fn status(&self) -> u16;
fn into_stream(self) -> Result<reqwest::Response, Self> { 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<Output = Error> + 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<u8> },
}
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<reqwest_middleware::Error> 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<Vec<u8>, 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)
}

/// Configure application transport policy once, including sending and response reads.
pub trait Transport: Sync {
type Error: From<Error> + Send;
fn execute<R: Response>(&self, request: reqwest_middleware::RequestBuilder, limit: usize)
-> impl std::future::Future<Output = Result<R, Self::Error>> + Send;
}
impl Transport for () {
type Error = Error;
async fn execute<R: Response>(&self, request: reqwest_middleware::RequestBuilder, limit: usize) -> Result<R, Error> {
R::decode(request.send().await.map_err(Error::from)?, limit).await
}
}

#[derive(Clone, Debug)]
pub struct Client {
pub struct Client<T = ()> {
http: reqwest_middleware::ClientWithMiddleware,
base_url: String,
bearer_token: Option<String>,
body_limit: usize,
transport: T,
}
fn encode_path(value: &str) -> String {
value.bytes().map(|byte| {
Expand All @@ -455,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<String>) -> 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<T: Transport> Client<T> {
/// 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<U: Transport>(self, transport: U) -> Client<U> {
Client { http: self.http, base_url: self.base_url, bearer_token: self.bearer_token, body_limit: self.body_limit, transport }
}
`
2 changes: 2 additions & 0 deletions moon.yml
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down
28 changes: 28 additions & 0 deletions testdata/fixtures/rust-responses.yaml
Original file line number Diff line number Diff line change
@@ -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
Loading