use kcode_speech_classification::{
Error as ClassifierError, ProtocolError, ProtocolRequest, ProtocolResponse, ProtocolResult,
SpeechClassifier,
};
use std::env;
use std::ffi::{OsStr, OsString};
use std::fs;
use std::io::{self, BufRead, BufReader, BufWriter, Read, Write};
use std::os::unix::fs::{FileTypeExt, MetadataExt, PermissionsExt};
use std::os::unix::net::{UnixListener, UnixStream};
use std::path::{Path, PathBuf};
use std::sync::{
Arc, Mutex,
mpsc::{self, Receiver, SyncSender, TrySendError},
};
use std::thread;
use std::time::Duration;
const MAX_CLIENT_THREADS: usize = 8;
const MAX_PENDING_CLIENTS: usize = 32;
const MAX_REQUEST_BYTES: usize = 1024 * 1024;
const CLIENT_READ_TIMEOUT: Duration = Duration::from_secs(30);
const CLIENT_WRITE_TIMEOUT: Duration = Duration::from_secs(5);
const BUSY_RESPONSE_TIMEOUT: Duration = Duration::from_millis(200);
struct Config {
database: PathBuf,
socket: PathBuf,
}
fn main() {
if let Err(error) = run() {
eprintln!("kcode-speech-classification: {error}");
std::process::exit(1);
}
}
fn run() -> Result<(), String> {
let config = parse_args(env::args_os().skip(1))?;
let classifier = SpeechClassifier::open(&config.database)
.map_err(|error| format!("could not open database: {error}"))?;
let listener = bind_socket(&config.socket)
.map_err(|error| format!("could not bind socket {}: {error}", config.socket.display()))?;
serve(listener, Arc::new(classifier)).map_err(|error| format!("socket service failed: {error}"))
}
fn parse_args(arguments: impl IntoIterator<Item = OsString>) -> Result<Config, String> {
let mut arguments = arguments.into_iter();
let mut database = None;
let mut socket = None;
while let Some(flag) = arguments.next() {
if flag == OsStr::new("--database") {
if database.is_some() {
return Err("--database may be supplied only once".to_owned());
}
database = Some(PathBuf::from(
arguments
.next()
.ok_or_else(|| "--database requires a path".to_owned())?,
));
} else if flag == OsStr::new("--socket") {
if socket.is_some() {
return Err("--socket may be supplied only once".to_owned());
}
socket = Some(PathBuf::from(
arguments
.next()
.ok_or_else(|| "--socket requires a path".to_owned())?,
));
} else {
return Err(format!("unknown argument {}", flag.to_string_lossy()));
}
}
let database = database.ok_or_else(|| "--database is required".to_owned())?;
let socket = socket.ok_or_else(|| "--socket is required".to_owned())?;
if database.as_os_str().is_empty() {
return Err("--database path must not be empty".to_owned());
}
if socket.as_os_str().is_empty() {
return Err("--socket path must not be empty".to_owned());
}
Ok(Config { database, socket })
}
fn serve(listener: UnixListener, classifier: Arc<SpeechClassifier>) -> io::Result<()> {
let (sender, receiver) = mpsc::sync_channel(MAX_PENDING_CLIENTS);
let receiver = Arc::new(Mutex::new(receiver));
for index in 0..MAX_CLIENT_THREADS {
let classifier = Arc::clone(&classifier);
let receiver = Arc::clone(&receiver);
let worker = thread::Builder::new()
.name(format!("speech-classifier-{index}"))
.spawn(move || worker_loop(&receiver, &classifier))?;
drop(worker);
}
for incoming in listener.incoming() {
match incoming {
Ok(stream) => dispatch_client(&sender, stream)?,
Err(error) if error.kind() == io::ErrorKind::Interrupted => {}
Err(error) => return Err(error),
}
}
Ok(())
}
fn dispatch_client(sender: &SyncSender<UnixStream>, stream: UnixStream) -> io::Result<()> {
match sender.try_send(stream) {
Ok(()) => Ok(()),
Err(TrySendError::Full(mut stream)) => {
stream.set_write_timeout(Some(BUSY_RESPONSE_TIMEOUT))?;
let response = ProtocolResponse::Error {
error: ProtocolError {
code: "server_busy".to_owned(),
message: "the bounded client queue is full".to_owned(),
},
};
write_response(&mut stream, &response)
}
Err(TrySendError::Disconnected(_)) => Err(io::Error::new(
io::ErrorKind::BrokenPipe,
"all client workers stopped",
)),
}
}
fn worker_loop(receiver: &Mutex<Receiver<UnixStream>>, classifier: &SpeechClassifier) {
loop {
let stream = {
let Ok(receiver) = receiver.lock() else {
return;
};
let Ok(stream) = receiver.recv() else {
return;
};
stream
};
if let Err(error) = handle_client(stream, classifier) {
eprintln!("kcode-speech-classification: client error: {error}");
}
}
}
fn handle_client(stream: UnixStream, classifier: &SpeechClassifier) -> io::Result<()> {
stream.set_read_timeout(Some(CLIENT_READ_TIMEOUT))?;
stream.set_write_timeout(Some(CLIENT_WRITE_TIMEOUT))?;
let reader_stream = stream.try_clone()?;
let mut reader = BufReader::new(reader_stream);
let mut writer = BufWriter::new(stream);
loop {
let mut line = Vec::new();
let bytes_read = {
let mut limited = Read::by_ref(&mut reader).take((MAX_REQUEST_BYTES + 1) as u64);
limited.read_until(b'\n', &mut line)?
};
if bytes_read == 0 {
return Ok(());
}
if bytes_read > MAX_REQUEST_BYTES {
let response = protocol_error(
"request_too_large",
"request line exceeds the configured byte limit",
);
write_response(&mut writer, &response)?;
return Ok(());
}
if line.last() != Some(&b'\n') {
let response = protocol_error("invalid_request", "request must end with a newline");
write_response(&mut writer, &response)?;
return Ok(());
}
line.pop();
if line.last() == Some(&b'\r') {
line.pop();
}
let response = match serde_json::from_slice::<ProtocolRequest>(&line) {
Ok(request) => execute_request(classifier, request),
Err(error) => protocol_error("invalid_request", &error.to_string()),
};
write_response(&mut writer, &response)?;
}
}
fn execute_request(classifier: &SpeechClassifier, request: ProtocolRequest) -> ProtocolResponse {
let result = match request {
ProtocolRequest::Identify {
key,
cohort,
row,
threshold,
} => classifier
.identify(key, cohort, row, threshold)
.map(ProtocolResult::Identify),
ProtocolRequest::Train {
key,
cohort,
row,
speaker_id,
} => classifier
.train(key, cohort, row, speaker_id)
.map(ProtocolResult::Train),
ProtocolRequest::Delete { key } => classifier.delete(key).map(ProtocolResult::Delete),
};
match result {
Ok(result) => ProtocolResponse::Success { result },
Err(error) => {
let code = classifier_error_code(&error).to_owned();
ProtocolResponse::Error {
error: ProtocolError {
code,
message: error.to_string(),
},
}
}
}
}
fn classifier_error_code(error: &ClassifierError) -> &'static str {
match error {
ClassifierError::Validation { .. } => "validation",
ClassifierError::Conflict { .. } => "conflict",
ClassifierError::UnsupportedSchema { .. } => "unsupported_schema",
ClassifierError::Storage(_) => "storage",
ClassifierError::CorruptStorage(_) => "corrupt_storage",
}
}
fn protocol_error(code: &str, message: &str) -> ProtocolResponse {
ProtocolResponse::Error {
error: ProtocolError {
code: code.to_owned(),
message: message.to_owned(),
},
}
}
fn write_response(writer: &mut impl Write, response: &ProtocolResponse) -> io::Result<()> {
serde_json::to_writer(&mut *writer, response).map_err(io::Error::other)?;
writer.write_all(b"\n")?;
writer.flush()
}
fn bind_socket(path: &Path) -> io::Result<UnixListener> {
prepare_socket_path(path)?;
let listener = UnixListener::bind(path)?;
if let Err(error) = fs::set_permissions(path, fs::Permissions::from_mode(0o600)) {
drop(listener);
let _ = fs::remove_file(path);
return Err(error);
}
Ok(listener)
}
fn prepare_socket_path(path: &Path) -> io::Result<()> {
let original = match fs::symlink_metadata(path) {
Ok(metadata) => metadata,
Err(error) if error.kind() == io::ErrorKind::NotFound => return Ok(()),
Err(error) => return Err(error),
};
if !original.file_type().is_socket() {
return Err(io::Error::new(
io::ErrorKind::AlreadyExists,
"refusing to remove an existing non-socket path",
));
}
match UnixStream::connect(path) {
Ok(_) => Err(io::Error::new(
io::ErrorKind::AddrInUse,
"an active server is already listening on the socket",
)),
Err(error) if error.kind() == io::ErrorKind::NotFound => Ok(()),
Err(error) if error.kind() == io::ErrorKind::ConnectionRefused => {
let current = match fs::symlink_metadata(path) {
Ok(metadata) => metadata,
Err(error) if error.kind() == io::ErrorKind::NotFound => return Ok(()),
Err(error) => return Err(error),
};
if !current.file_type().is_socket()
|| current.dev() != original.dev()
|| current.ino() != original.ino()
{
return Err(io::Error::new(
io::ErrorKind::AlreadyExists,
"socket path changed while stale state was being checked",
));
}
match fs::remove_file(path) {
Ok(()) => Ok(()),
Err(error) if error.kind() == io::ErrorKind::NotFound => Ok(()),
Err(error) => Err(error),
}
}
Err(error) => Err(io::Error::new(
error.kind(),
format!("could not establish whether the socket is active: {error}"),
)),
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::sync::atomic::{AtomicU64, Ordering};
static NEXT_PATH: AtomicU64 = AtomicU64::new(0);
fn socket_path(label: &str) -> PathBuf {
env::temp_dir().join(format!(
"kcode-speech-classification-bin-{}-{label}-{}.sock",
std::process::id(),
NEXT_PATH.fetch_add(1, Ordering::Relaxed)
))
}
#[test]
fn arguments_require_exact_database_and_socket_flags() {
let config = parse_args([
OsString::from("--socket"),
OsString::from("/tmp/service.sock"),
OsString::from("--database"),
OsString::from("/tmp/service.sqlite3"),
])
.unwrap();
assert_eq!(config.database, Path::new("/tmp/service.sqlite3"));
assert_eq!(config.socket, Path::new("/tmp/service.sock"));
assert!(parse_args([OsString::from("--database")]).is_err());
assert!(
parse_args([
OsString::from("--database"),
OsString::from("db"),
OsString::from("--unknown"),
OsString::from("value"),
])
.is_err()
);
}
#[test]
fn refuses_and_preserves_a_non_socket_path() {
let path = socket_path("ordinary-file");
fs::write(&path, b"preserve me").unwrap();
let error = prepare_socket_path(&path).unwrap_err();
assert_eq!(error.kind(), io::ErrorKind::AlreadyExists);
assert_eq!(fs::read(&path).unwrap(), b"preserve me");
fs::remove_file(path).unwrap();
}
#[test]
fn distinguishes_live_and_stale_socket_paths() {
let live_path = socket_path("live");
let live_listener = UnixListener::bind(&live_path).unwrap();
let error = prepare_socket_path(&live_path).unwrap_err();
assert_eq!(error.kind(), io::ErrorKind::AddrInUse);
assert!(
fs::symlink_metadata(&live_path)
.unwrap()
.file_type()
.is_socket()
);
drop(live_listener);
fs::remove_file(live_path).unwrap();
let stale_path = socket_path("stale");
let stale_listener = UnixListener::bind(&stale_path).unwrap();
drop(stale_listener);
prepare_socket_path(&stale_path).unwrap();
assert!(!stale_path.exists());
}
#[test]
fn bound_socket_is_owner_only() {
let path = socket_path("permissions");
let listener = bind_socket(&path).unwrap();
let mode = fs::metadata(&path).unwrap().permissions().mode() & 0o777;
assert_eq!(mode, 0o600);
drop(listener);
fs::remove_file(path).unwrap();
}
}