#![cfg_attr(
all(target_os = "windows", not(debug_assertions)),
windows_subsystem = "windows"
)]
use std::{
collections::HashMap,
sync::{Arc, Mutex},
time::Duration,
};
use clap::{CommandFactory, FromArgMatches, Parser};
use dbx_tools_core::{build_info, init_logging_with_verbose};
use dbx_tools_service::{
tray::{
replace_menu, run_tray, CheckMenuItem, Icon, Menu, MenuEvent, MenuId, MenuItem,
PredefinedMenuItem, Submenu, TrayCli, TrayConfig, TrayControl, TrayHandle, TrayIcon,
},
ServiceConfig,
};
use reqwest::blocking::Client;
use serde::{Deserialize, Serialize};
const CONTROL_HEADER: &str = "x-model-proxy-control";
#[derive(Clone, Debug, Parser)]
#[command(name = "dbx-model-proxy-tray")]
struct Cli {
#[arg(long, default_value = "http://127.0.0.1:4000")]
url: String,
#[command(flatten)]
tray: TrayCli,
}
#[derive(Clone, Debug, Deserialize)]
#[serde(rename_all = "camelCase")]
struct RuntimeStatus {
profile: String,
}
#[derive(Clone, Debug, Deserialize)]
#[serde(rename_all = "camelCase")]
struct AuthStatus {
runtime: RuntimeStatus,
}
#[derive(Clone, Debug, Deserialize)]
struct ProfileSummary {
name: String,
}
#[derive(Clone, Debug, Deserialize)]
struct Profiles {
profiles: Vec<ProfileSummary>,
}
#[derive(Clone, Debug, Serialize)]
#[serde(rename_all = "camelCase", tag = "kind", content = "profile")]
enum RuntimeSelection {
Profile(String),
}
#[derive(Clone)]
struct ApiClient {
base_url: String,
client: Client,
}
impl ApiClient {
fn new(base_url: String) -> Result<Self, Box<dyn std::error::Error>> {
Ok(Self {
base_url: base_url.trim_end_matches('/').to_owned(),
client: Client::builder().timeout(Duration::from_secs(20)).build()?,
})
}
fn status(&self) -> Result<AuthStatus, Box<dyn std::error::Error>> {
Ok(self
.client
.get(format!("{}/api/auth", self.base_url))
.send()?
.error_for_status()?
.json()?)
}
fn wait_for_status(&self) -> AuthStatus {
loop {
if let Ok(status) = self.status() {
return status;
}
std::thread::sleep(Duration::from_millis(500));
}
}
fn profiles(&self) -> Result<Vec<String>, Box<dyn std::error::Error>> {
let profiles = self
.client
.get(format!("{}/api/auth/profiles", self.base_url))
.send()?
.error_for_status()?
.json::<Profiles>()?
.profiles;
Ok(profiles.into_iter().map(|profile| profile.name).collect())
}
fn switch(&self, profile: String) -> Result<AuthStatus, Box<dyn std::error::Error>> {
Ok(self
.client
.put(format!("{}/api/auth", self.base_url))
.header(reqwest::header::ORIGIN, &self.base_url)
.header(CONTROL_HEADER, "1")
.json(&RuntimeSelection::Profile(profile))
.send()?
.error_for_status()?
.json()?)
}
}
#[derive(Clone)]
enum Action {
Open(String),
Profile(String),
Quit,
}
fn proxy_targets(url: &str) -> [(&'static str, String); 3] {
let url = url.trim_end_matches('/');
[
("Open Models", format!("{url}/v1/models")),
("Open API", format!("{url}/api")),
("Open GraphQL", format!("{url}/graphql")),
]
}
fn menu(
url: &str,
current: &str,
profiles: Option<&[String]>,
) -> Result<(Menu, HashMap<MenuId, Action>), Box<dyn std::error::Error>> {
let menu = Menu::new();
let address = Submenu::new(format!("Proxy: {url}"), true);
let profile_menu = Submenu::new(current, true);
let mut actions = HashMap::new();
for (label, target) in proxy_targets(url) {
let item = MenuItem::new(label, true, None);
actions.insert(item.id().clone(), Action::Open(target));
address.append(&item)?;
}
if let Some(profiles) = profiles {
for profile in profiles {
let item = CheckMenuItem::new(profile, true, profile == current, None);
actions.insert(item.id().clone(), Action::Profile(profile.clone()));
profile_menu.append(&item)?;
}
} else {
profile_menu.append(&MenuItem::new("Loading profiles...", false, None))?;
}
let quit = MenuItem::new("Quit", true, None);
actions.insert(quit.id().clone(), Action::Quit);
menu.append_items(&[
&address,
&profile_menu,
&PredefinedMenuItem::separator(),
&quit,
])?;
Ok((menu, actions))
}
fn refresh_profiles(
handle: TrayHandle,
api: ApiClient,
current: Arc<Mutex<String>>,
actions: Arc<Mutex<HashMap<MenuId, Action>>>,
) {
std::thread::spawn(move || match api.profiles() {
Ok(profiles) => {
let profile = current
.lock()
.map(|profile| profile.clone())
.unwrap_or_else(|_| "Profile unavailable".to_owned());
let url = api.base_url.clone();
handle.dispatch(move |tray| match menu(&url, &profile, Some(&profiles)) {
Ok((menu, next_actions)) => {
if let Ok(mut registered) = actions.lock() {
*registered = next_actions;
}
if let Err(error) = replace_menu(tray, menu) {
tracing::error!(%error, "tray profile menu update failed");
}
}
Err(error) => tracing::error!(%error, "tray profile menu could not be built"),
});
}
Err(error) => tracing::warn!(%error, "Databricks profiles could not be loaded"),
});
}
fn proxy_icon() -> Result<Icon, Box<dyn std::error::Error>> {
let mut rgba = Vec::with_capacity(32 * 32 * 4);
let color = if cfg!(target_os = "macos") {
[0x00, 0x00, 0x00, 0xff]
} else {
[0xff, 0x36, 0x21, 0xff]
};
let rectangles = [
(4, 13, 6, 6),
(24, 4, 6, 6),
(24, 13, 6, 6),
(24, 22, 6, 6),
(10, 15, 6, 2),
(16, 6, 2, 20),
(18, 6, 6, 2),
(18, 15, 6, 2),
(18, 24, 6, 2),
];
for y in 0..32 {
for x in 0..32 {
let inside = rectangles.iter().any(|(left, top, width, height)| {
(*left..*left + *width).contains(&x) && (*top..*top + *height).contains(&y)
});
rgba.extend_from_slice(if inside { &color } else { &[0, 0, 0, 0] });
}
}
Ok(Icon::from_rgba(rgba, 32, 32)?)
}
fn main() -> Result<(), Box<dyn std::error::Error>> {
let cli = Cli::from_arg_matches(&Cli::command().version(build_info::version()).get_matches())?;
init_logging_with_verbose(false)?;
let api = ApiClient::new(cli.url.clone())?;
let current_profile = if cli.tray.probe {
"Model Proxy".to_owned()
} else {
api.wait_for_status().runtime.profile
};
let current = Arc::new(Mutex::new(current_profile.clone()));
let actions = Arc::new(Mutex::new(HashMap::<MenuId, Action>::new()));
let handle = Arc::new(Mutex::new(None::<TrayHandle>));
let (initial_menu, initial_actions) = menu(&cli.url, ¤t_profile, None)?;
*actions.lock().map_err(|_| "tray action lock is poisoned")? = initial_actions;
let ready_api = api.clone();
let ready_current = Arc::clone(¤t);
let ready_actions = Arc::clone(&actions);
let ready_handle = Arc::clone(&handle);
let on_ready = Arc::new(move |tray_handle: TrayHandle| {
if let Ok(mut stored) = ready_handle.lock() {
*stored = Some(tray_handle.clone());
}
refresh_profiles(
tray_handle,
ready_api.clone(),
Arc::clone(&ready_current),
Arc::clone(&ready_actions),
);
});
let menu_api = api.clone();
let menu_current = Arc::clone(¤t);
let menu_actions = Arc::clone(&actions);
let menu_handle = Arc::clone(&handle);
let on_menu = Arc::new(move |event: &MenuEvent, _tray: &TrayIcon| {
let action = menu_actions
.lock()
.map_err(|_| "tray action lock is poisoned")?
.get(&event.id)
.cloned();
match action {
Some(Action::Quit) => Ok(TrayControl::Quit),
Some(Action::Open(target)) => {
std::thread::spawn(move || {
if let Err(error) = open::that(&target) {
tracing::error!(%error, url = target, "tray link could not be opened");
}
});
Ok(TrayControl::Continue)
}
Some(Action::Profile(profile)) => {
let Some(tray_handle) = menu_handle
.lock()
.map_err(|_| "tray handle lock is poisoned")?
.clone()
else {
return Ok(TrayControl::Continue);
};
let api = menu_api.clone();
let current = Arc::clone(&menu_current);
let actions = Arc::clone(&menu_actions);
std::thread::spawn(move || match api.switch(profile) {
Ok(status) => {
if let Ok(mut selected) = current.lock() {
*selected = status.runtime.profile;
}
refresh_profiles(tray_handle, api, current, actions);
}
Err(error) => tracing::error!(%error, "Databricks profile switch failed"),
});
Ok(TrayControl::Continue)
}
None => Ok(TrayControl::Continue),
}
});
let service = ServiceConfig::new("model-proxy", 4000)?
.with_invalid_runtime_detector(dbx_tools_core::is_databricks_app);
run_tray(
cli.tray,
TrayConfig {
service,
title: "Model Proxy".to_owned(),
icon: proxy_icon()?,
menu: initial_menu,
icon_as_template: cfg!(target_os = "macos"),
on_ready,
on_menu,
},
)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn proxy_targets_use_the_configured_base_url() {
assert_eq!(
proxy_targets("http://127.0.0.1:4001/"),
[
("Open Models", "http://127.0.0.1:4001/v1/models".to_owned()),
("Open API", "http://127.0.0.1:4001/api".to_owned()),
("Open GraphQL", "http://127.0.0.1:4001/graphql".to_owned()),
]
);
}
}