Skip to main content

aiway_plugin/
wasm_ctx.rs

1//! WASM 侧 HttpContext 实现
2//!
3//! 通过宿主函数访问网关的真实请求数据,而非在 WASM 内部维护独立状态。
4//! 所有数据读取均委托给宿主侧的 `aiway::host_xxx` 函数。
5
6use crate::PluginError;
7use crate::plugin_ctx::{HttpRequest, HttpResponse, PluginContext};
8#[cfg(feature = "model")]
9use aiway_protocol::model::Provider;
10use bytes::Bytes;
11use http::Uri;
12use std::any::Any;
13// ---------------------------------------------------------------------------
14// 宿主函数 FFI 声明
15// ---------------------------------------------------------------------------
16
17#[link(wasm_import_module = "aiway")]
18unsafe extern "C" {
19    fn host_request_id(buf_ptr: *mut u8, buf_len: i32) -> i32;
20    fn host_request_ts() -> i64;
21    fn host_is_sse() -> i32;
22    fn host_is_websocket() -> i32;
23    fn host_get_request_header(
24        name_ptr: *const u8,
25        name_len: i32,
26        buf_ptr: *mut u8,
27        buf_len: i32,
28    ) -> i32;
29    fn host_get_response_header(
30        name_ptr: *const u8,
31        name_len: i32,
32        buf_ptr: *mut u8,
33        buf_len: i32,
34    ) -> i32;
35    fn host_method(buf_ptr: *mut u8, buf_len: i32) -> i32;
36    fn host_uri(buf_ptr: *mut u8, buf_len: i32) -> i32;
37    fn host_set_uri(uri_ptr: *const u8, uri_len: i32);
38    fn host_status() -> i32;
39    fn host_get_route_name(buf_ptr: *mut u8, buf_len: i32) -> i32;
40    fn host_get_routing_url(buf_ptr: *mut u8, buf_len: i32) -> i32;
41    fn host_get_response_body_size() -> i64;
42    fn host_set_response_body_size(size: i64);
43    fn host_log(level: i32, msg_ptr: *const u8, msg_len: i32);
44    #[cfg(feature = "model")]
45    fn host_get_model_name(buf_ptr: *mut u8, buf_len: i32) -> i32;
46    #[cfg(feature = "model")]
47    fn host_get_model_provider(buf_ptr: *mut u8, buf_len: i32) -> i32;
48    fn host_set_request_header(
49        name_ptr: *const u8,
50        name_len: i32,
51        value_ptr: *const u8,
52        value_len: i32,
53    );
54    fn host_set_response_header(
55        name_ptr: *const u8,
56        name_len: i32,
57        value_ptr: *const u8,
58        value_len: i32,
59    );
60    fn host_append_request_header(
61        name_ptr: *const u8,
62        name_len: i32,
63        value_ptr: *const u8,
64        value_len: i32,
65    );
66    fn host_append_response_header(
67        name_ptr: *const u8,
68        name_len: i32,
69        value_ptr: *const u8,
70        value_len: i32,
71    );
72    fn host_remove_request_header(name_ptr: *const u8, name_len: i32);
73    fn host_remove_response_header(name_ptr: *const u8, name_len: i32);
74    fn host_http_request(
75        req_ptr: *const u8,
76        req_len: i32,
77        resp_buf_ptr: *mut u8,
78        resp_buf_len: i32,
79    ) -> i32;
80    fn host_config(buf_ptr: *mut u8, buf_len: i32) -> i32;
81    fn host_get_request_body(buf_ptr: *mut u8, buf_len: i32) -> i32;
82    fn host_set_request_body(body_ptr: *const u8, body_len: i32);
83    fn host_get_response_body(buf_ptr: *mut u8, buf_len: i32) -> i32;
84    fn host_set_response_body(body_ptr: *const u8, body_len: i32);
85    fn host_respond(
86        status: i32,
87        hdr_ptr: *const u8,
88        hdr_len: i32,
89        body_ptr: *const u8,
90        body_len: i32,
91    );
92}
93
94// ---------------------------------------------------------------------------
95// WasmHttpContext
96// ---------------------------------------------------------------------------
97
98/// WASM 侧的插件上下文实现。
99///
100/// 不持有任何状态,所有数据通过宿主函数按需获取。
101pub struct WasmHttpContext;
102
103/// 通过宿主函数读取字符串。
104///
105/// `f` 为宿主函数,遵循 snprintf 语义:返回数据实际长度(可能大于 `buf_len`)。
106/// `initial_len` 为初始缓冲区大小,若数据超出则自动扩容重试。
107fn read_host_string(
108    f: unsafe extern "C" fn(*mut u8, i32) -> i32,
109    initial_len: i32,
110) -> Option<String> {
111    read_host_bytes(f, initial_len).map(|bytes| String::from_utf8_lossy(&bytes).to_string())
112}
113
114/// 通过宿主函数读取字节数组。
115///
116/// 语义同 [`read_host_string`]:初始缓冲区不足时自动扩容重试。
117fn read_host_bytes(
118    f: unsafe extern "C" fn(*mut u8, i32) -> i32,
119    initial_len: i32,
120) -> Option<Vec<u8>> {
121    let mut buf = vec![0u8; initial_len as usize];
122    let needed = unsafe { f(buf.as_mut_ptr(), initial_len) };
123    if needed <= 0 {
124        return None;
125    }
126    let needed = needed as usize;
127    if needed > buf.len() {
128        buf.resize(needed, 0);
129        let len = unsafe { f(buf.as_mut_ptr(), needed as i32) };
130        if len <= 0 {
131            return None;
132        }
133        return Some(buf[..len as usize].to_vec());
134    }
135    Some(buf[..needed].to_vec())
136}
137
138/// 通知宿主记录插件主动响应(由 SDK 导出宏在 `Outcome::Respond` 时调用)。
139///
140/// headers 以 bincode 序列化传入,Host 侧反序列化后存入 HttpContext。
141pub fn respond_to_host(status: u16, headers: Vec<(String, String)>, body: Vec<u8>) {
142    let headers_bytes = bincode::serialize(&headers).unwrap_or_default();
143    unsafe {
144        host_respond(
145            status as i32,
146            headers_bytes.as_ptr(),
147            headers_bytes.len() as i32,
148            body.as_ptr(),
149            body.len() as i32,
150        );
151    }
152}
153
154/// 通过宿主函数读取 bincode 序列化数据并反序列化。
155///
156/// 语义同 [`read_host_string`]:初始缓冲区不足时自动扩容重试。
157#[cfg(feature = "model")]
158fn read_host_bincode<T: serde::de::DeserializeOwned>(
159    f: unsafe extern "C" fn(*mut u8, i32) -> i32,
160    initial_len: i32,
161) -> Option<T> {
162    let mut buf = vec![0u8; initial_len as usize];
163    let needed = unsafe { f(buf.as_mut_ptr(), initial_len) };
164    if needed <= 0 {
165        return None;
166    }
167    let needed = needed as usize;
168    if needed > buf.len() {
169        buf.resize(needed, 0);
170        let len = unsafe { f(buf.as_mut_ptr(), needed as i32) };
171        if len <= 0 {
172            return None;
173        }
174        return bincode::deserialize(&buf[..len as usize]).ok();
175    }
176    bincode::deserialize(&buf[..needed]).ok()
177}
178
179impl PluginContext for WasmHttpContext {
180    fn request_id(&self) -> String {
181        read_host_string(host_request_id, 64).unwrap_or_default()
182    }
183
184    fn request_ts(&self) -> i64 {
185        unsafe { host_request_ts() }
186    }
187
188    fn is_sse(&self) -> bool {
189        unsafe { host_is_sse() != 0 }
190    }
191
192    fn is_websocket(&self) -> bool {
193        unsafe { host_is_websocket() != 0 }
194    }
195
196    fn get_request_header(&self, name: &str) -> Option<String> {
197        let name_bytes = name.as_bytes();
198        let mut buf = vec![0u8; 256];
199        let needed = unsafe {
200            host_get_request_header(
201                name_bytes.as_ptr(),
202                name_bytes.len() as i32,
203                buf.as_mut_ptr(),
204                buf.len() as i32,
205            )
206        };
207        if needed <= 0 {
208            return None;
209        }
210        let needed = needed as usize;
211        if needed > buf.len() {
212            buf.resize(needed, 0);
213            let len = unsafe {
214                host_get_request_header(
215                    name_bytes.as_ptr(),
216                    name_bytes.len() as i32,
217                    buf.as_mut_ptr(),
218                    buf.len() as i32,
219                )
220            };
221            if len <= 0 {
222                return None;
223            }
224            return Some(String::from_utf8_lossy(&buf[..len as usize]).to_string());
225        }
226        Some(String::from_utf8_lossy(&buf[..needed]).to_string())
227    }
228
229    fn get_response_header(&self, name: &str) -> Option<String> {
230        let name_bytes = name.as_bytes();
231        let mut buf = vec![0u8; 256];
232        let needed = unsafe {
233            host_get_response_header(
234                name_bytes.as_ptr(),
235                name_bytes.len() as i32,
236                buf.as_mut_ptr(),
237                buf.len() as i32,
238            )
239        };
240        if needed <= 0 {
241            return None;
242        }
243        let needed = needed as usize;
244        if needed > buf.len() {
245            buf.resize(needed, 0);
246            let len = unsafe {
247                host_get_response_header(
248                    name_bytes.as_ptr(),
249                    name_bytes.len() as i32,
250                    buf.as_mut_ptr(),
251                    buf.len() as i32,
252                )
253            };
254            if len <= 0 {
255                return None;
256            }
257            return Some(String::from_utf8_lossy(&buf[..len as usize]).to_string());
258        }
259        Some(String::from_utf8_lossy(&buf[..needed]).to_string())
260    }
261
262    fn method(&self) -> Option<String> {
263        read_host_string(host_method, 16)
264    }
265
266    fn uri(&self) -> Option<Uri> {
267        read_host_string(host_uri, 512).and_then(|s| s.parse().ok())
268    }
269
270    fn set_uri(&mut self, uri: Uri) {
271        let bytes = uri.to_string();
272        let b = bytes.as_bytes();
273        unsafe { host_set_uri(b.as_ptr(), b.len() as i32) }
274    }
275
276    fn status(&self) -> Option<u16> {
277        let v = unsafe { host_status() };
278        if v < 0 { None } else { Some(v as u16) }
279    }
280
281    fn get_route_name(&self) -> Option<String> {
282        read_host_string(host_get_route_name, 256)
283    }
284
285    fn get_routing_url(&self) -> Option<String> {
286        read_host_string(host_get_routing_url, 512)
287    }
288
289    fn get_response_body_size(&self) -> Option<i64> {
290        let v = unsafe { host_get_response_body_size() };
291        if v < 0 { None } else { Some(v) }
292    }
293
294    fn set_response_body_size(&mut self, size: i64) {
295        unsafe { host_set_response_body_size(size) }
296    }
297
298    fn set_request_header(&mut self, name: &str, value: &str) {
299        let nb = name.as_bytes();
300        let vb = value.as_bytes();
301        unsafe {
302            host_set_request_header(nb.as_ptr(), nb.len() as i32, vb.as_ptr(), vb.len() as i32)
303        }
304    }
305
306    fn set_response_header(&mut self, name: &str, value: &str) {
307        let nb = name.as_bytes();
308        let vb = value.as_bytes();
309        unsafe {
310            host_set_response_header(nb.as_ptr(), nb.len() as i32, vb.as_ptr(), vb.len() as i32)
311        }
312    }
313
314    fn append_request_header(&mut self, name: &str, value: &str) {
315        let nb = name.as_bytes();
316        let vb = value.as_bytes();
317        unsafe {
318            host_append_request_header(nb.as_ptr(), nb.len() as i32, vb.as_ptr(), vb.len() as i32)
319        }
320    }
321
322    fn append_response_header(&mut self, name: &str, value: &str) {
323        let nb = name.as_bytes();
324        let vb = value.as_bytes();
325        unsafe {
326            host_append_response_header(nb.as_ptr(), nb.len() as i32, vb.as_ptr(), vb.len() as i32)
327        }
328    }
329
330    fn remove_request_header(&mut self, name: &str) {
331        let nb = name.as_bytes();
332        unsafe { host_remove_request_header(nb.as_ptr(), nb.len() as i32) }
333    }
334
335    fn remove_response_header(&mut self, name: &str) {
336        let nb = name.as_bytes();
337        unsafe { host_remove_response_header(nb.as_ptr(), nb.len() as i32) }
338    }
339
340    #[cfg(feature = "model")]
341    fn get_model_name(&self) -> Option<String> {
342        read_host_string(host_get_model_name, 256)
343    }
344
345    #[cfg(feature = "model")]
346    fn get_model_provider(&self) -> Option<Provider> {
347        read_host_bincode(host_get_model_provider, 512)
348    }
349
350    fn log(&self, level: i32, msg: &str) {
351        let bytes = msg.as_bytes();
352        unsafe { host_log(level, bytes.as_ptr(), bytes.len() as i32) }
353    }
354
355    fn http_request(&self, req: &HttpRequest) -> Result<HttpResponse, PluginError> {
356        let req_bytes = bincode::serialize(req)
357            .map_err(|e| PluginError::HttpError(format!("serialize request failed: {e}")))?;
358
359        // 初始缓冲区 4KB,不足时扩容重试
360        let mut buf = vec![0u8; 4096];
361        loop {
362            let needed = unsafe {
363                host_http_request(
364                    req_bytes.as_ptr(),
365                    req_bytes.len() as i32,
366                    buf.as_mut_ptr(),
367                    buf.len() as i32,
368                )
369            };
370            if needed < 0 {
371                return Err(PluginError::HttpError(format!(
372                    "http_request failed with code {needed}"
373                )));
374            }
375            if needed == 0 {
376                return Err(PluginError::HttpError(
377                    "http_request returned empty response".into(),
378                ));
379            }
380            let needed = needed as usize;
381            if needed > buf.len() {
382                buf.resize(needed, 0);
383                continue;
384            }
385            return bincode::deserialize(&buf[..needed])
386                .map_err(|e| PluginError::HttpError(format!("deserialize response failed: {e}")));
387        }
388    }
389
390    fn config(&self) -> Option<String> {
391        read_host_string(host_config, 1024)
392    }
393
394    fn request_body(&self) -> Option<Bytes> {
395        read_host_bytes(host_get_request_body, 4096).map(Bytes::from)
396    }
397
398    fn set_request_body(&mut self, body: Vec<u8>) {
399        unsafe { host_set_request_body(body.as_ptr(), body.len() as i32) }
400    }
401
402    fn response_body(&self) -> Option<Bytes> {
403        read_host_bytes(host_get_response_body, 4096).map(Bytes::from)
404    }
405
406    fn set_response_body(&mut self, body: Vec<u8>) {
407        unsafe { host_set_response_body(body.as_ptr(), body.len() as i32) }
408    }
409
410    fn as_any_mut(&mut self) -> &mut dyn Any {
411        self
412    }
413}