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:

  1. Import Necessary Libraries: We import libraries like url, hyper, hyper_proxy, gjson, and tokio to handle HTTP requests, routing, and JSON parsing.
  2. Create the Reverse Proxy:
    • The new_openai_reverse_proxy function sets up the proxy using hyper_proxy. It specifies Intercept::All to intercept all requests. You'll need to replace "https://some.azurewebsites.net" with your actual Azure OpenAI endpoint.
  3. Handle Incoming Requests:
    • The handle function receives an incoming request (req) and the proxy instance.
    • It parses the request's headers, method, and body.
    • The gjson crate is used to extract the model from 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-version query 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.
  4. Main Function:
    • The main function 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 handle function to process each request.
    • The server continues running until it encounters an error or is explicitly stopped.

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.

Go to Rust:  Rewriting an OpenAI Reverse Proxy

原文地址: https://www.cveoy.top/t/topic/lQUl 著作权归作者所有。请勿转载和采集!

免费AI点我,无需注册和登录