use std::sync::Arc;
use axum::body::Bytes;
use axum::extract::{Path, State};
use axum::http::header::{CACHE_CONTROL, CONTENT_TYPE};
use axum::http::{HeaderMap, HeaderValue, StatusCode};
use axum::response::{IntoResponse, Response};
use axum::routing::{get, post, put};
use axum::{Extension, Json, Router};
use kcode_k1_access_full_audio::{
AccessId, FullAudioState, GroupId, K1AccessFullAudio, PersonId, SpeakerLabelV1, UserId,
};
use kcode_k1_access_profiles::{ProfileId, ProfileSelection, TxId};
use kcode_k1_http::Principal;
use serde::{Deserialize, Serialize};
use serde_json::Value;
pub fn authenticated_routes(audio: Arc<K1AccessFullAudio>) -> Router<()> {
Router::new()
.route("/audio", get(list_audio))
.route("/audio/profiles/{profile_id}", post(submit_audio))
.route("/audio/groups/{group_id}", get(list_group_audio))
.route("/audio/{access_id}", get(get_audio))
.route(
"/audio/{access_id}/fragments/{fragment_id}/audio",
get(get_fragment_audio),
)
.route(
"/audio/{access_id}/fragments/{fragment_id}/labels",
put(submit_labels),
)
.route(
"/audio/{access_id}/fragments/{fragment_id}/retry",
post(retry_fragment),
)
.route(
"/audio/{access_id}/fragments/{fragment_id}/discard",
post(discard_fragment),
)
.with_state(audio)
}
#[derive(Debug)]
struct ApiError {
status: StatusCode,
code: &'static str,
message: String,
}
type ApiResult<T> = Result<T, ApiError>;
#[derive(Serialize)]
struct ErrorDto {
error: &'static str,
message: String,
}
impl IntoResponse for ApiError {
fn into_response(self) -> Response {
json_response(
self.status,
ErrorDto {
error: self.code,
message: self.message,
},
)
}
}
fn problem(status: StatusCode, code: &'static str, message: impl Into<String>) -> ApiError {
ApiError {
status,
code,
message: message.into(),
}
}
fn invalid_id(code: &'static str, subject: &str) -> ApiError {
problem(
StatusCode::BAD_REQUEST,
code,
format!("{subject} ID is not canonical lowercase hexadecimal"),
)
}
fn dependency_error(operation: &str, message: String) -> ApiError {
if message.contains("fragment does not belong to full audio") {
return problem(
StatusCode::NOT_FOUND,
"fragment_not_found",
"fragment is unavailable to the caller",
);
}
if message.contains("cannot access full audio") || message == "access denied" {
return problem(
StatusCode::NOT_FOUND,
"audio_not_found",
"audio is unavailable to the caller",
);
}
problem(
StatusCode::SERVICE_UNAVAILABLE,
"audio_unavailable",
format!("{operation}: {message}"),
)
}
fn json_response<T: Serialize>(status: StatusCode, value: T) -> Response {
let mut response = (status, Json(value)).into_response();
response
.headers_mut()
.insert(CACHE_CONTROL, HeaderValue::from_static("no-store"));
response
}
fn no_content() -> Response {
let mut response = StatusCode::NO_CONTENT.into_response();
response
.headers_mut()
.insert(CACHE_CONTROL, HeaderValue::from_static("no-store"));
response
}
async fn run<T, F>(operation: &'static str, task: F) -> ApiResult<T>
where
T: Send + 'static,
F: FnOnce() -> Result<T, String> + Send + 'static,
{
match tokio::task::spawn_blocking(task).await {
Ok(Ok(value)) => Ok(value),
Ok(Err(message)) => Err(dependency_error(operation, message)),
Err(_) => Err(problem(
StatusCode::SERVICE_UNAVAILABLE,
"audio_unavailable",
format!("{operation}: blocking task failed"),
)),
}
}
fn caller(principal: &Principal) -> UserId {
UserId::from_tx_id(TxId::from_bytes(*principal.user_id()))
}
fn parse_txid(value: &str, code: &'static str, subject: &str) -> ApiResult<TxId> {
if value.len() != 24 {
return Err(invalid_id(code, subject));
}
let mut output = [0; 12];
for (index, pair) in value.as_bytes().chunks_exact(2).enumerate() {
let digit = |byte| match byte {
b'0'..=b'9' => Some(byte - b'0'),
b'a'..=b'f' => Some(byte - b'a' + 10),
_ => None,
};
output[index] = digit(pair[0])
.zip(digit(pair[1]))
.map(|(high, low)| high << 4 | low)
.ok_or_else(|| invalid_id(code, subject))?;
}
Ok(TxId::from_bytes(output))
}
fn parse_access_id(value: &str) -> ApiResult<AccessId> {
parse_txid(value, "invalid_access_id", "access").map(AccessId::new)
}
fn parse_fragment_id(value: &str) -> ApiResult<TxId> {
parse_txid(value, "invalid_fragment_id", "fragment")
}
fn parse_group_id(value: &str) -> ApiResult<GroupId> {
parse_txid(value, "invalid_group_id", "group").map(GroupId::new)
}
fn parse_profile_id(value: &str) -> ApiResult<ProfileId> {
parse_txid(value, "invalid_profile_id", "profile").map(ProfileId::new)
}
fn hex(bytes: &[u8; 12]) -> String {
const DIGITS: &[u8; 16] = b"0123456789abcdef";
let mut value = String::with_capacity(24);
for byte in bytes {
value.push(DIGITS[(byte >> 4) as usize] as char);
value.push(DIGITS[(byte & 15) as usize] as char);
}
value
}
fn access_hex(id: &AccessId) -> String {
hex(id.txid().as_bytes())
}
fn require_content_type(headers: &HeaderMap, expected: &str) -> ApiResult<()> {
let mut values = headers.get_all(CONTENT_TYPE).iter();
let valid = values
.next()
.and_then(|value| value.to_str().ok())
.is_some_and(|value| {
value
.split(';')
.next()
.is_some_and(|media| media.trim().eq_ignore_ascii_case(expected))
})
&& values.next().is_none();
if valid {
Ok(())
} else {
Err(problem(
StatusCode::UNSUPPORTED_MEDIA_TYPE,
"unsupported_media_type",
"request content type is unsupported",
))
}
}
#[derive(Serialize)]
struct SubmittedDto {
access_id: String,
full_audio_id: String,
}
#[derive(Serialize)]
struct FragmentDto {
fragment_id: String,
start_sample_48k: u64,
end_sample_48k: u64,
status: Value,
}
#[derive(Serialize)]
struct StatusDto {
state: &'static str,
fragments: Vec<FragmentDto>,
final_transcript: Option<String>,
}
fn state_name(state: FullAudioState) -> &'static str {
match state {
FullAudioState::Processing => "processing",
FullAudioState::AwaitingLabels => "awaiting_labels",
FullAudioState::NeedsAttention => "needs_attention",
FullAudioState::Complete => "complete",
}
}
#[derive(Deserialize)]
struct LabelsDto {
labels: Vec<LabelDto>,
}
#[derive(Deserialize)]
struct LabelDto {
speaker: String,
person_id: Option<String>,
}
fn parse_labels(input: LabelsDto) -> ApiResult<Vec<SpeakerLabelV1>> {
input
.labels
.into_iter()
.map(|label| {
let person_id = label
.person_id
.map(|value| {
parse_txid(&value, "invalid_person_id", "person").map(PersonId::from_tx_id)
})
.transpose()?;
Ok(SpeakerLabelV1 {
speaker: label.speaker.parse().map_err(|_| {
problem(
StatusCode::BAD_REQUEST,
"invalid_labels",
"speaker label is invalid",
)
})?,
person_id,
})
})
.collect()
}
async fn submit_audio(
State(audio): State<Arc<K1AccessFullAudio>>,
Extension(principal): Extension<Principal>,
Path(value): Path<String>,
headers: HeaderMap,
body: Bytes,
) -> ApiResult<Response> {
require_content_type(&headers, "application/octet-stream")?;
let profile = ProfileSelection::Saved(parse_profile_id(&value)?);
let user = caller(&principal);
let submitted = run("submit audio", move || {
audio.submit_for_user(user, profile, &body)
})
.await?;
Ok(json_response(
StatusCode::CREATED,
SubmittedDto {
access_id: access_hex(&submitted.access_id),
full_audio_id: hex(submitted.full_audio_id.as_bytes()),
},
))
}
async fn list_audio(
State(audio): State<Arc<K1AccessFullAudio>>,
Extension(principal): Extension<Principal>,
) -> ApiResult<Response> {
let user = caller(&principal);
let ids = run("list audio", move || audio.list_for_user(user)).await?;
Ok(json_response(
StatusCode::OK,
ids.iter().map(access_hex).collect::<Vec<_>>(),
))
}
async fn list_group_audio(
State(audio): State<Arc<K1AccessFullAudio>>,
Extension(principal): Extension<Principal>,
Path(value): Path<String>,
) -> ApiResult<Response> {
let group = parse_group_id(&value)?;
let user = caller(&principal);
let ids = run("list group audio", move || {
audio.list_group_for_user(user, group)
})
.await?;
Ok(json_response(
StatusCode::OK,
ids.iter().map(access_hex).collect::<Vec<_>>(),
))
}
async fn get_audio(
State(audio): State<Arc<K1AccessFullAudio>>,
Extension(principal): Extension<Principal>,
Path(value): Path<String>,
) -> ApiResult<Response> {
let access_id = parse_access_id(&value)?;
let user = caller(&principal);
let status = run("get audio status", move || {
audio.status_for_user(user, access_id)
})
.await?;
let state = state_name(status.state);
let mut fragments = Vec::with_capacity(status.fragments.len());
for fragment in status.fragments {
let serialized = serde_json::to_value(fragment.status).map_err(|_| {
problem(
StatusCode::SERVICE_UNAVAILABLE,
"audio_unavailable",
"serialize audio status: status serialization failed",
)
})?;
fragments.push(FragmentDto {
fragment_id: hex(fragment.fragment_id.as_bytes()),
start_sample_48k: fragment.start_sample_48k,
end_sample_48k: fragment.end_sample_48k,
status: serialized,
});
}
Ok(json_response(
StatusCode::OK,
StatusDto {
state,
fragments,
final_transcript: status.final_transcript,
},
))
}
async fn get_fragment_audio(
State(audio): State<Arc<K1AccessFullAudio>>,
Extension(principal): Extension<Principal>,
Path((access_value, fragment_value)): Path<(String, String)>,
) -> ApiResult<Response> {
let access_id = parse_access_id(&access_value)?;
let fragment_id = parse_fragment_id(&fragment_value)?;
let user = caller(&principal);
let bytes = run("get fragment audio", move || {
audio.fragment_audio_for_user(user, access_id, fragment_id)
})
.await?;
Ok((
[
(CONTENT_TYPE, HeaderValue::from_static("audio/ogg")),
(CACHE_CONTROL, HeaderValue::from_static("no-store")),
],
bytes,
)
.into_response())
}
async fn submit_labels(
State(audio): State<Arc<K1AccessFullAudio>>,
Extension(principal): Extension<Principal>,
Path((access_value, fragment_value)): Path<(String, String)>,
headers: HeaderMap,
body: Bytes,
) -> ApiResult<Response> {
require_content_type(&headers, "application/json")?;
let input: LabelsDto = serde_json::from_slice(&body).map_err(|_| {
problem(
StatusCode::BAD_REQUEST,
"invalid_json",
"request body is invalid",
)
})?;
let labels = parse_labels(input)?;
let access_id = parse_access_id(&access_value)?;
let fragment_id = parse_fragment_id(&fragment_value)?;
let user = caller(&principal);
run("submit fragment labels", move || {
audio.submit_labels_for_user(user, access_id, fragment_id, labels)
})
.await?;
Ok(no_content())
}
async fn retry_fragment(
State(audio): State<Arc<K1AccessFullAudio>>,
Extension(principal): Extension<Principal>,
Path((access_value, fragment_value)): Path<(String, String)>,
) -> ApiResult<Response> {
let access_id = parse_access_id(&access_value)?;
let fragment_id = parse_fragment_id(&fragment_value)?;
let user = caller(&principal);
run("retry fragment", move || {
audio.retry_fragment_for_user(user, access_id, fragment_id)
})
.await?;
Ok(no_content())
}
async fn discard_fragment(
State(audio): State<Arc<K1AccessFullAudio>>,
Extension(principal): Extension<Principal>,
Path((access_value, fragment_value)): Path<(String, String)>,
) -> ApiResult<Response> {
let access_id = parse_access_id(&access_value)?;
let fragment_id = parse_fragment_id(&fragment_value)?;
let user = caller(&principal);
run("discard fragment", move || {
audio.discard_fragment_for_user(user, access_id, fragment_id)
})
.await?;
Ok(no_content())
}
#[cfg(test)]
mod tests {
use super::*;
use serde_json::json;
#[test]
fn canonical_ids_and_user_conversion_preserve_bytes() {
let id = "00112233445566778899aabb";
assert_eq!(hex(parse_fragment_id(id).unwrap().as_bytes()), id);
assert!(parse_fragment_id("00112233445566778899AABB").is_err());
assert!(parse_fragment_id("0011").is_err());
assert_eq!(
hex(caller_bytes([1; 12]).as_tx_id().as_bytes()),
"010101010101010101010101"
);
}
fn caller_bytes(bytes: [u8; 12]) -> UserId {
UserId::from_tx_id(TxId::from_bytes(bytes))
}
#[test]
fn labels_accept_known_and_unknown_people() {
let labels = parse_labels(LabelsDto {
labels: vec![
LabelDto {
speaker: "Speaker 1".to_owned(),
person_id: None,
},
LabelDto {
speaker: "Speaker 2".to_owned(),
person_id: Some("00112233445566778899aabb".to_owned()),
},
],
})
.unwrap();
assert_eq!(labels.len(), 2);
assert!(labels[0].person_id.is_none());
assert_eq!(
hex(labels[1].person_id.unwrap().as_tx_id().as_bytes()),
"00112233445566778899aabb"
);
assert!(
parse_labels(LabelsDto {
labels: vec![LabelDto {
speaker: "Unknown".to_owned(),
person_id: None
}]
})
.is_err()
);
}
#[test]
fn status_shape_and_hidden_errors_are_stable() {
let value = serde_json::to_value(StatusDto {
state: "processing",
fragments: Vec::new(),
final_transcript: None,
})
.unwrap();
assert_eq!(
value,
json!({"state":"processing","fragments":[],"final_transcript":null})
);
let error = dependency_error(
"get audio status",
"principal cannot access full audio".to_owned(),
);
assert_eq!(error.status, StatusCode::NOT_FOUND);
assert_eq!(error.code, "audio_not_found");
assert!(!error.message.contains("principal"));
}
#[test]
fn content_types_and_public_route_assembler_are_exact() {
let mut headers = HeaderMap::new();
headers.insert(
CONTENT_TYPE,
HeaderValue::from_static("application/json; charset=utf-8"),
);
assert!(require_content_type(&headers, "application/json").is_ok());
assert!(require_content_type(&headers, "application/octet-stream").is_err());
let _ = authenticated_routes as fn(Arc<K1AccessFullAudio>) -> Router<()>;
}
}