1use std::sync::Arc;
10
11use sim_kernel::{
12 Consistency, Cx, Error, EvalFabric, EvalMode, EvalReply, EvalRequest, Expr, Result, Symbol,
13 eval_remote_capability,
14};
15use sim_value::{access, build};
16
17use crate::{FabricEvalSite, ServerAddress};
18
19pub const MIC_CAPTURE_NAMESPACE: &str = "watch";
21
22pub const MIC_CAPTURE_KIND: &str = "mic-capture";
24
25pub const XR_MIC_CHUNK_NAMESPACE: &str = "xr";
27
28pub const XR_MIC_CHUNK_KIND: &str = "mic-chunk";
30
31pub const ASR_TRANSCRIPT_NAMESPACE: &str = "asr";
33
34pub const ASR_TRANSCRIPT_KIND: &str = "transcript";
36
37const MODELED_ASR_SITE_KIND: &str = "watch-asr";
38const MODELED_GLASSES_ASR_SITE_KIND: &str = "glasses-asr";
39
40#[derive(Clone, Debug)]
42pub struct ModeledAsrFabric {
43 label: String,
44}
45
46impl ModeledAsrFabric {
47 pub fn new(label: impl Into<String>) -> Self {
49 Self {
50 label: label.into(),
51 }
52 }
53}
54
55impl EvalFabric for ModeledAsrFabric {
56 fn realize(&self, cx: &mut Cx, request: EvalRequest) -> Result<EvalReply> {
57 if matches!(request.consistency, Consistency::RemoteOnly) {
58 return Err(Error::CapabilityDenied {
59 capability: eval_remote_capability(),
60 });
61 }
62 if !matches!(request.mode, EvalMode::Eval) {
63 return Err(Error::Eval(
64 "modeled ASR fabric only supports eval mode".to_owned(),
65 ));
66 }
67 for capability in &request.required_capabilities {
68 cx.require(capability)?;
69 }
70
71 let input = AsrInput::from_expr(&request.expr)?;
72 let trace = input.trace_symbol();
73 let output = match input {
74 AsrInput::Watch(capture) => transcript_expr(
75 &self.label,
76 capture.seq,
77 capture.frame_count,
78 capture.byte_count,
79 ),
80 AsrInput::Glasses(chunk) => voice_intent_expr(&chunk),
81 };
82 Ok(EvalReply {
83 value: cx.factory().expr(output)?,
84 diagnostics: cx.take_diagnostics(),
85 trace: request
86 .trace
87 .then(|| cx.factory().symbol(trace))
88 .transpose()?,
89 })
90 }
91}
92
93pub fn modeled_asr_site(
95 address: ServerAddress,
96 codecs: Vec<Symbol>,
97 label: impl Into<String>,
98) -> FabricEvalSite {
99 FabricEvalSite::new(
100 MODELED_ASR_SITE_KIND,
101 address,
102 codecs,
103 Arc::new(ModeledAsrFabric::new(label)),
104 )
105}
106
107pub fn modeled_glasses_asr_site(
109 address: ServerAddress,
110 codecs: Vec<Symbol>,
111 label: impl Into<String>,
112) -> FabricEvalSite {
113 FabricEvalSite::new(
114 MODELED_GLASSES_ASR_SITE_KIND,
115 address,
116 codecs,
117 Arc::new(ModeledAsrFabric::new(label)),
118 )
119}
120
121#[derive(Clone, Debug, PartialEq, Eq)]
122enum AsrInput {
123 Watch(MicCaptureView),
124 Glasses(XrMicChunkView),
125}
126
127impl AsrInput {
128 fn from_expr(expr: &Expr) -> Result<Self> {
129 let Some(kind) = access::field_sym(expr, "kind") else {
130 return Err(Error::HostError(
131 "expected watch mic capture or xr mic chunk".to_owned(),
132 ));
133 };
134 if kind.namespace.as_deref() == Some(MIC_CAPTURE_NAMESPACE)
135 && kind.name.as_ref() == MIC_CAPTURE_KIND
136 {
137 return MicCaptureView::from_expr(expr).map(Self::Watch);
138 }
139 if kind.namespace.as_deref() == Some(XR_MIC_CHUNK_NAMESPACE)
140 && kind.name.as_ref() == XR_MIC_CHUNK_KIND
141 {
142 return XrMicChunkView::from_expr(expr).map(Self::Glasses);
143 }
144 Err(Error::HostError(
145 "expected watch mic capture or xr mic chunk".to_owned(),
146 ))
147 }
148
149 fn trace_symbol(&self) -> Symbol {
150 match self {
151 Self::Watch(_) => Symbol::qualified("asr", "modeled-watch"),
152 Self::Glasses(_) => Symbol::qualified("asr", "modeled-glasses"),
153 }
154 }
155}
156
157#[derive(Clone, Copy, Debug, PartialEq, Eq)]
158struct MicCaptureView {
159 seq: u64,
160 frame_count: usize,
161 byte_count: usize,
162}
163
164impl MicCaptureView {
165 fn from_expr(expr: &Expr) -> Result<Self> {
166 ensure_kind(
167 expr,
168 MIC_CAPTURE_NAMESPACE,
169 MIC_CAPTURE_KIND,
170 "watch mic capture",
171 )?;
172 ensure_no_extra(
173 expr,
174 &["kind", "frames", "seq", "sample-rate-hz", "channels"],
175 "watch mic capture",
176 )?;
177 let frames = match access::required(expr, "frames", "watch mic capture")? {
178 Expr::List(items) if !items.is_empty() => items,
179 Expr::List(_) => {
180 return Err(Error::HostError(
181 "watch mic capture requires at least one frame".to_owned(),
182 ));
183 }
184 _ => {
185 return Err(Error::TypeMismatch {
186 expected: "audio frame list",
187 found: "non-list",
188 });
189 }
190 };
191 let mut byte_count = 0usize;
192 for frame in frames {
193 byte_count += frame_byte_count(frame)?;
194 }
195 Ok(Self {
196 seq: uint_field(expr, "seq", "watch mic capture")?,
197 frame_count: frames.len(),
198 byte_count,
199 })
200 }
201}
202
203#[derive(Clone, Debug, PartialEq, Eq)]
204struct XrMicChunkView {
205 ref_id: Symbol,
206 seq: u64,
207 byte_count: u64,
208}
209
210impl XrMicChunkView {
211 fn from_expr(expr: &Expr) -> Result<Self> {
212 ensure_kind(
213 expr,
214 XR_MIC_CHUNK_NAMESPACE,
215 XR_MIC_CHUNK_KIND,
216 "xr mic chunk",
217 )?;
218 ensure_no_extra(
219 expr,
220 &["kind", "ref", "seq", "sample-rate-hz", "channels", "bytes"],
221 "xr mic chunk",
222 )?;
223 let ref_id = match access::required(expr, "ref", "xr mic chunk")? {
224 Expr::Symbol(symbol) => symbol.clone(),
225 _ => {
226 return Err(Error::TypeMismatch {
227 expected: "audio chunk reference symbol",
228 found: "non-symbol",
229 });
230 }
231 };
232 let sample_rate_hz = uint_field(expr, "sample-rate-hz", "xr mic chunk")?;
233 let channels = uint_field(expr, "channels", "xr mic chunk")?;
234 let byte_count = uint_field(expr, "bytes", "xr mic chunk")?;
235 if sample_rate_hz == 0 || channels == 0 || byte_count == 0 {
236 return Err(Error::HostError(
237 "xr mic chunk requires nonzero audio metadata".to_owned(),
238 ));
239 }
240 Ok(Self {
241 ref_id,
242 seq: uint_field(expr, "seq", "xr mic chunk")?,
243 byte_count,
244 })
245 }
246}
247
248fn frame_byte_count(expr: &Expr) -> Result<usize> {
249 ensure_kind(
250 expr,
251 MIC_CAPTURE_NAMESPACE,
252 "audio-frame",
253 "watch audio frame",
254 )?;
255 ensure_no_extra(expr, &["kind", "at-ms", "pcm"], "watch audio frame")?;
256 match access::required(expr, "pcm", "watch audio frame")? {
257 Expr::Bytes(bytes) => Ok(bytes.len()),
258 _ => Err(Error::TypeMismatch {
259 expected: "raw PCM bytes",
260 found: "non-bytes",
261 }),
262 }
263}
264
265fn transcript_expr(label: &str, seq: u64, frame_count: usize, byte_count: usize) -> Expr {
266 build::map(vec![
267 (
268 "kind",
269 build::qsym(ASR_TRANSCRIPT_NAMESPACE, ASR_TRANSCRIPT_KIND),
270 ),
271 (
272 "text",
273 build::text(format!(
274 "{label}: seq {seq}, {frame_count} frame(s), {byte_count} byte(s)"
275 )),
276 ),
277 ("seq", build::uint(seq)),
278 ("frames", build::uint(frame_count as u64)),
279 ("bytes", build::uint(byte_count as u64)),
280 ])
281}
282
283fn voice_intent_expr(chunk: &XrMicChunkView) -> Expr {
284 build::map(vec![
285 ("kind", build::qsym("intent", "invoke")),
286 (
287 "origin",
288 build::map(vec![
289 ("operator", build::sym("agent")),
290 ("at-tick", build::uint(chunk.seq)),
291 ]),
292 ),
293 ("target", build::sym("focused")),
294 (
295 "op",
296 Expr::Symbol(Symbol::qualified("glasses/voice", "modeled-asr")),
297 ),
298 (
299 "args",
300 build::list(vec![
301 Expr::Symbol(chunk.ref_id.clone()),
302 build::map(vec![("bytes", build::uint(chunk.byte_count))]),
303 ]),
304 ),
305 ])
306}
307
308fn ensure_kind(expr: &Expr, namespace: &str, kind: &str, context: &str) -> Result<()> {
309 match access::field_sym(expr, "kind") {
310 Some(symbol)
311 if symbol.namespace.as_deref() == Some(namespace) && symbol.name.as_ref() == kind =>
312 {
313 Ok(())
314 }
315 _ => Err(Error::HostError(format!("expected {context}"))),
316 }
317}
318
319fn ensure_no_extra(expr: &Expr, allowed: &[&str], context: &str) -> Result<()> {
320 let Expr::Map(entries) = expr else {
321 return Err(Error::HostError(format!("expected {context}")));
322 };
323 for (key, _) in entries {
324 let Expr::Symbol(symbol) = key else {
325 return Err(Error::HostError(format!(
326 "{context} has a non-symbol field"
327 )));
328 };
329 if symbol.namespace.is_some() || !allowed.contains(&symbol.name.as_ref()) {
330 return Err(Error::HostError(format!(
331 "{context} has unexpected field {}",
332 symbol.as_qualified_str()
333 )));
334 }
335 }
336 Ok(())
337}
338
339fn uint_field(expr: &Expr, name: &str, context: &str) -> Result<u64> {
340 match access::required(expr, name, context)? {
341 Expr::Number(number) if number.domain.namespace.is_none() => number
342 .canonical
343 .parse()
344 .map_err(|_| Error::Eval(format!("{context} field {name} is not u64"))),
345 _ => Err(Error::Eval(format!("{context} field {name} is not u64"))),
346 }
347}