kcode-speech-classification 0.1.0

Typed open-set speaker classification with SQLite-backed training and an NDJSON Unix-socket service.
Documentation
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();
    }
}