use crate::context::Context;
use crate::connection::ZitiStream;
use crate::error::{ZitiError, ZitiResult};
use crate::service::list_services;
use crate::transport::http::controller_client;
use crate::transport::protocol::{ContentType, ZitiMessage};
use crate::transport::{TlsConfig, WebSocketTransport};
use bytes::Bytes;
use tokio_tungstenite::tungstenite::Message;
use url::Url;
pub async fn dial(service_name: &str, context: &Context) -> ZitiResult<ZitiStream> {
let services = list_services(context.session_manager()).await?;
let service = services
.iter()
.find(|s| s.name == service_name)
.ok_or_else(|| ZitiError::ServiceNotFound {
service_name: service_name.to_string(),
})?;
let edge_routers = get_service_terminators(&service.id, context).await?;
if edge_routers.is_empty() {
return Err(ZitiError::ConnectionFailed(
format!("No edge routers available for service '{}'", service_name)
));
}
let edge_router = &edge_routers[0];
let tls_config = TlsConfig::from_identity(context.identity_manager())?;
let ws_url = Url::parse(&format!("wss://{}:{}/ws", edge_router.hostname, edge_router.port))
.map_err(|e| ZitiError::ConfigError(format!("Invalid edge router URL: {}", e)))?;
let mut transport = WebSocketTransport::connect(ws_url, tls_config).await?;
perform_ziti_handshake(&mut transport, &service.id).await?;
Ok(ZitiStream::from_transport(transport))
}
#[derive(Debug, Clone, serde::Deserialize)]
pub struct EdgeRouter {
pub hostname: String,
pub port: u16,
#[serde(default)]
pub supported_protocols: Vec<String>,
}
#[derive(Debug, serde::Deserialize)]
struct Terminator {
#[allow(dead_code)]
pub id: String,
pub router: Option<EdgeRouter>,
}
#[derive(Debug, serde::Deserialize)]
struct TerminatorsResponse {
data: Vec<Terminator>,
}
async fn get_service_terminators(service_id: &str, context: &Context) -> ZitiResult<Vec<EdgeRouter>> {
let api_session = context.session_manager().get_api_session().await?;
let client = controller_client(context.identity_manager(), context.connect_timeout()).await?;
let terminators_url = format!(
"{}/services/{}/terminators",
context.identity_manager().zt_api().trim_end_matches('/'),
service_id
);
let response = client
.get(&terminators_url)
.header("Content-Type", "application/json")
.header("zt-session", &api_session.token)
.send()
.await
.map_err(|e| {
ZitiError::ConnectionFailed(format!("Failed to get service terminators: {}", e))
})?;
if !response.status().is_success() {
let status = response.status();
let error_text = response
.text()
.await
.unwrap_or_else(|_| "Unknown error".to_string());
return Err(ZitiError::ProtocolError {
message: format!(
"Terminators request failed with status {}: {}",
status, error_text
),
});
}
let terminators_response: TerminatorsResponse = response
.json()
.await
.map_err(|e| ZitiError::ProtocolError {
message: format!("Failed to parse terminators response: {}", e),
})?;
let edge_routers = terminators_response
.data
.into_iter()
.filter_map(|t| t.router)
.collect();
Ok(edge_routers)
}
async fn perform_ziti_handshake(
transport: &mut WebSocketTransport,
service_id: &str,
) -> ZitiResult<()> {
let hello_bytes = build_hello_message(service_id)?;
transport.send(Message::Binary(hello_bytes)).await?;
match transport.receive().await? {
Some(Message::Binary(data)) => parse_hello_response(&data),
Some(_) => Err(ZitiError::ProtocolError {
message: "Unexpected message type in handshake response".to_string(),
}),
None => Err(ZitiError::ConnectionFailed(
"Connection closed during handshake".to_string(),
)),
}
}
fn build_hello_message(service_id: &str) -> ZitiResult<Bytes> {
let mut hello = ZitiMessage::new(ContentType::Hello, 1, Bytes::new());
hello.header
.add_header("service_id".to_string(), service_id.to_string());
hello.header
.add_header("version".to_string(), "1.0".to_string());
hello.serialize()
}
fn parse_hello_response(data: &[u8]) -> ZitiResult<()> {
let msg = ZitiMessage::deserialize(Bytes::copy_from_slice(data))?;
if msg.content_type() != ContentType::Hello {
return Err(ZitiError::ProtocolError {
message: format!(
"Expected Hello response, got content_type={:?}",
msg.content_type()
),
});
}
crate::connection::listen::check_response_status(msg, "hello")
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_edge_router_deserialization() {
let json = r#"{
"hostname": "edge-router.example.com",
"port": 443,
"supported_protocols": ["tls", "ws"]
}"#;
let edge_router: EdgeRouter = serde_json::from_str(json).unwrap();
assert_eq!(edge_router.hostname, "edge-router.example.com");
assert_eq!(edge_router.port, 443);
assert_eq!(edge_router.supported_protocols, vec!["tls", "ws"]);
}
#[test]
fn test_hello_message_roundtrip() {
let bytes = build_hello_message("test-service").unwrap();
let parsed = ZitiMessage::deserialize(bytes).unwrap();
assert_eq!(parsed.content_type(), ContentType::Hello);
assert_eq!(parsed.header.get_header("service_id").unwrap(), "test-service");
assert_eq!(parsed.header.get_header("version").unwrap(), "1.0");
}
#[test]
fn test_parse_hello_response_ok() {
let mut msg = ZitiMessage::new(ContentType::Hello, 2, Bytes::new());
msg.header
.add_header("status".to_string(), "ok".to_string());
let bytes = msg.serialize().unwrap();
assert!(parse_hello_response(&bytes).is_ok());
}
#[test]
fn test_parse_hello_response_error_carries_detail() {
let mut msg = ZitiMessage::new(ContentType::Hello, 2, Bytes::new());
msg.header
.add_header("status".to_string(), "error".to_string());
msg.header
.add_header("error".to_string(), "policy denied".to_string());
let bytes = msg.serialize().unwrap();
let err = parse_hello_response(&bytes).unwrap_err();
let s = format!("{}", err);
assert!(s.contains("policy denied"), "got: {}", s);
}
#[test]
fn test_parse_hello_response_wrong_type() {
let msg = ZitiMessage::new(ContentType::Data, 1, Bytes::from_static(b"x"));
let bytes = msg.serialize().unwrap();
assert!(parse_hello_response(&bytes).is_err());
}
}