Go to Rust: Rewriting an OpenAI Reverse Proxy
use url::Url;
use std::io::{self, Read};
use std::path;
use std::fmt;
use std::collections::HashMap;
use std::sync::Arc;
use hyper::{Body, Request, Response, Server, Uri};
use hyper::client::HttpConnector;
use hyper::header::{HeaderMap, HeaderValue, AUTHORIZATION};
use hyper::http::request::Builder;
use hyper::http::Method;
use hyper::service::{make_service_fn, service_fn};
use hyper::server::conn::AddrStream;
use hyper_proxy::{Intercept, Proxy, ProxyConnector};
use tokio::runtime::Runtime;
use futures_util::StreamExt;
fn new_openai_reverse_proxy() -> Arc<Proxy<HttpConnector, Intercept>> {
let azure_openai_endpoint = "https://some.azurewebsites.net";
let remote = Url::parse(azure_openai_endpoint).unwrap();
let mut proxies = HashMap::new();
proxies.insert(remote.scheme().to_owned(), remote.host_str().unwrap().to_owned());
let proxy = Proxy::new(Intercept::All, ProxyConnector::new(proxies).unwrap());
Arc::new(proxy)
}
async fn handle(req: Request<Body>, proxy: Arc<Proxy<HttpConnector, Intercept>>) -> Result<Response<Body>, hyper::Error> {
let mut builder = Builder::new();
builder.method(req.method().clone()).uri(req.uri().clone());
let mut headers = HeaderMap::new();
for (key, value) in req.headers() {
headers.insert(key.clone(), value.clone());
}
let mut body = Vec::new();
while let Some(chunk) = req.body_mut().next().await {
let data = chunk?;
body.extend(data);
}
let model = gjson::get(&body, 'model').unwrap().to_string();
let deployment = get_deployment_by_model(&model);
let token = headers.get(AUTHORIZATION).unwrap().to_str().unwrap().replace("Bearer ", "");
headers.insert('api-key', HeaderValue::from_str(&token).unwrap());
headers.remove(AUTHORIZATION);
let mut origin_url = String::new();
origin_url.push_str(req.uri().scheme().unwrap());
origin_url.push_str("://");
origin_url.push_str(req.uri().host().unwrap());
origin_url.push_str(req.uri().path());
let remote = format!("/{}/openai/deployments/{}", remote, deployment);
let uri = Uri::from_str(&remote).unwrap();
let mut query_pairs = uri.query_pairs().collect::<Vec<_>>();
query_pairs.push(('api-version'.to_owned(), AzureOpenAIAPIVersion.to_owned()));
let query_pairs = query_pairs.into_iter().map(|(k, v)| format!("{}={}", k, v)).collect::<Vec<_>>().join("&");
let new_uri = format!("{}?{}", uri.path(), query_pairs);
builder.uri(new_uri.as_str());
builder.headers_mut().unwrap().extend(headers);
let resp = proxy
.call(builder.body(Body::from(body)).unwrap())
.await
.unwrap();
if resp.headers().get("Content-Type").unwrap() == "text/event-stream" {
let (parts, body) = resp.into_parts();
let mut body_bytes = Vec::new();
body.into_bytes().await.unwrap().iter().for_each(|b| {
body_bytes.push(*b);
});
body_bytes.push(b'\n');
Response::from_parts(parts, Body::from(body_bytes))
} else {
Ok(resp)
}
}
fn main() {
let proxy = new_openai_reverse_proxy();
let addr = "127.0.0.1:3000".parse().unwrap();
let make_svc = make_service_fn(|socket: &AddrStream| {
let remote_addr = socket.remote_addr();
let proxy = proxy.clone();
async move {
Ok::<_, hyper::Error>(service_fn(move |req: Request<Body>| {
let proxy = proxy.clone();
handle(req, proxy).map(move |mut res| {
let mut headers = res.headers_mut();
headers.insert("access-control-allow-origin", HeaderValue::from_static("*"));
headers.insert("access-control-allow-methods", HeaderValue::from_static("GET, POST, PUT, DELETE, OPTIONS"));
headers.insert("access-control-allow-headers", HeaderValue::from_static("Content-Type, Authorization"));
res
})
}))
}
});
let mut rt = Runtime::new().unwrap();
rt.block_on(async {
let server = Server::bind(&addr).serve(make_svc);
println!("Listening on http://{}", addr);
server.await.unwrap();
});
}
Explanation:
- Import Necessary Libraries: We import libraries like
url,hyper,hyper_proxy,gjson, andtokioto handle HTTP requests, routing, and JSON parsing. - Create the Reverse Proxy:
- The
new_openai_reverse_proxyfunction sets up the proxy usinghyper_proxy. It specifiesIntercept::Allto intercept all requests. You'll need to replace"https://some.azurewebsites.net"with your actual Azure OpenAI endpoint.
- The
- Handle Incoming Requests:
- The
handlefunction receives an incoming request (req) and the proxy instance. - It parses the request's headers, method, and body.
- The
gjsoncrate is used to extract themodelfrom the request body. - It modifies the headers to include the
api-key(extracted from the original Authorization header) and removes the Authorization header. - The URL is constructed to point to the correct Azure OpenAI deployment based on the
model. - The
api-versionquery parameter is added. - The modified request is then forwarded to the Azure OpenAI service using the proxy.
- Bugfix: For requests with
Content-Type: text/event-stream, it appends a newline character to the response body to address a known difference between OpenAI and Azure responses.
- The
- Main Function:
- The
mainfunction starts the HTTP server, binds it to the specified address (127.0.0.1:3000), and listens for incoming requests. - It sets up a service that uses the
handlefunction to process each request. - The server continues running until it encounters an error or is explicitly stopped.
- The
This code provides a solid foundation for building an OpenAI reverse proxy using Rust. Remember to install the necessary dependencies and adjust the code according to your specific needs.
原文地址: https://www.cveoy.top/t/topic/lQUl 著作权归作者所有。请勿转载和采集!