use std::collections::HashMap;
use std::net::SocketAddr;
use flare_grpc_proto::capability::extension_plugin_server::{
ExtensionPlugin, ExtensionPluginServer,
};
use flare_grpc_proto::capability::{GenericRequest, GenericResponse};
use flare_plugin_host::{PluginDeclaration, PluginHost, SeatModel};
use tonic::{Request, Response, Status, transport::Server};
const PLUGIN_ID: &str = "echo-plugin";
const OPERATIONS: &[&str] = &["example.echo.say", "example.echo.upper"];
#[derive(Default)]
struct EchoPlugin;
#[tonic::async_trait]
impl ExtensionPlugin for EchoPlugin {
async fn call(
&self,
request: Request<GenericRequest>,
) -> Result<Response<GenericResponse>, Status> {
let req = request.into_inner();
let tenant = req.metadata.get("tenant_id").cloned().unwrap_or_default();
let user = req.metadata.get("user_id").cloned().unwrap_or_default();
let payload: serde_json::Value = match req.payload.as_ref() {
Some(any) => serde_json::from_slice(&any.value)
.map_err(|e| Status::invalid_argument(format!("payload 不是合法 JSON: {e}")))?,
None => serde_json::Value::Null,
};
let text = payload.get("text").and_then(|v| v.as_str()).unwrap_or("");
let result = match req.operation.as_str() {
"example.echo.say" => serde_json::json!({
"echo": text,
"tenant_id": tenant,
"user_id": user,
}),
"example.echo.upper" => serde_json::json!({ "echo": text.to_uppercase() }),
other => {
return Ok(Response::new(GenericResponse {
ok: false,
payload: None,
error_code: "UNKNOWN_OPERATION".into(),
error_message: format!("{other} 已声明但未实现"),
request_id: req.request_id,
}));
}
};
Ok(Response::new(GenericResponse {
ok: true,
payload: Some(prost_types::Any {
type_url: "type.googleapis.com/flare.capability.v1.PayloadJson".into(),
value: serde_json::to_vec(&result).map_err(|e| Status::internal(e.to_string()))?,
}),
error_code: String::new(),
error_message: String::new(),
request_id: req.request_id,
}))
}
}
#[tokio::main]
async fn main() -> Result<(), Box<dyn std::error::Error>> {
let listen: SocketAddr = std::env::var("ECHO_PLUGIN_LISTEN")
.unwrap_or_else(|_| "127.0.0.1:7901".into())
.parse()?;
let capability =
std::env::var("CAPABILITY_ENDPOINT").unwrap_or_else(|_| "http://127.0.0.1:50110".into());
let advertise = std::env::var("ECHO_PLUGIN_ADVERTISE").unwrap_or_else(|_| listen.to_string());
let declaration = PluginDeclaration {
tenant_id: "0".into(),
plugin_id: PLUGIN_ID.into(),
capability_id: OPERATIONS[0].into(),
grpc_authority: advertise.clone(),
plugin_version: env!("CARGO_PKG_VERSION").into(),
api_version: "1".into(),
manifest_sha256: String::new(),
declared_operations: OPERATIONS.iter().map(|s| s.to_string()).collect(),
labels: HashMap::new(),
seat_model: SeatModel::Tenant,
};
let server = tokio::spawn(
Server::builder()
.add_service(ExtensionPluginServer::new(EchoPlugin))
.serve_with_shutdown(listen, async {
tokio::signal::ctrl_c().await.ok();
}),
);
println!("echo-plugin 监听 {listen},对外通告 {advertise}");
match PluginHost::connect(&capability).await {
Ok(mut host) => {
host.announce(&declaration).await?;
println!("已向 {capability} 声明 {} 个 operation", OPERATIONS.len());
server.await??;
host.withdraw("0", PLUGIN_ID).await.ok();
println!("已注销");
}
Err(e) => {
println!("连不上 {capability}({e});服务继续,未注册");
server.await??;
}
}
Ok(())
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn declaration_matches_implementation() {
let implemented = include_str!("echo_plugin.rs");
for op in OPERATIONS {
assert!(
implemented.contains(&format!("\"{op}\" =>")),
"{op} 声明了却没有对应的分发臂"
);
}
}
}