Reconnect SRS after disconnect

This commit is contained in:
AviiNL
2023-12-23 09:18:44 +01:00
parent 05f4859f67
commit 945cf75242
2 changed files with 161 additions and 137 deletions

View File

@@ -56,6 +56,7 @@ fn connect_to_grpc(mut commands: Commands, tokio: Res<TokioResource>, url: Res<G
tokio::time::sleep(Duration::from_secs(5)).await;
continue;
};
println!("Connected to DCS: {}", url);
let Ok(mut stream) = client
.stream_units(StreamUnitsRequest {
@@ -86,10 +87,8 @@ fn connect_to_grpc(mut commands: Commands, tokio: Res<TokioResource>, url: Res<G
}
}
eprintln!("Disconnected");
eprintln!("Disconnected from DCS: {}", url);
tx.send(Response::Disconnected).ok();
tokio::time::sleep(Duration::from_secs(5)).await;
}
});

View File

@@ -164,142 +164,167 @@ fn listen_srs(
let voice = radio.voice;
let stt_url = stt_url.as_str().to_string();
radio.handle = Some(tokio.0.spawn(async move {
// todo: add loop here to reconnect on disconnect/errors
let tcp = TcpStream::connect(addr).await?;
let (tcp_stream, tcp_sink) = tcp.into_split();
let mut messages_sink = FramedWrite::new(tcp_sink, MessagesCodec::new());
let mut messages_stream = FramedRead::new(tcp_stream, MessagesCodec::new());
let udp = UdpSocket::bind(SocketAddr::from(([0, 0, 0, 0], 0))).await?;
udp.connect(addr).await?;
let mut voice_ping_interval = tokio::time::interval(Duration::from_secs(15));
let (mut voice_sink, mut voice_stream) = UdpFramed::new(udp, VoiceCodec::new()).split();
let mut packet_id = 1;
messages_sink
.send(MessageRequest::Sync(SyncMessageRequest {
msg_type: MsgType,
client,
version: SRS_VERSION.to_string(),
}))
.await?;
let transmissions = Arc::new(RwLock::new(HashMap::<u32, Vec<VoicePacket>>::new()));
let frames: FrameQueue<Vec<Vec<u8>>> = FrameQueue::new();
loop {
tokio::select! {
Some(data) = frames.next() => {
let start = tokio::time::Instant::now();
for (i, frame) in data.into_iter().enumerate() {
if frame.is_empty() {
continue;
}
let packet = VoicePacket {
audio_part: frame,
wav_audio_part: None,
frequencies: vec![Frequency {
freq: frequency as f64,
modulation: if frequency <= 87_995_000 {
voice_codec::Modulation::Fm
} else {
voice_codec::Modulation::Am
},
encryption: Encryption::None,
}],
unit_id: id.into(),
packet_id,
hop_count: 0,
transmission_sguid: sguid,
client_sguid: sguid,
};
voice_sink.send((packet.into(), addr)).await.ok();
packet_id = packet_id.wrapping_add(1);
let playtime = Duration::from_millis((i as u64 + 1) * 20); // 20m per frame count
let elapsed = start.elapsed();
if playtime > elapsed {
let s = playtime - elapsed;
sleep(s).await;
}
}
}
Some(data) = message_handle.recv() => {
messages_sink.send(data).await?;
}
Some(Ok(data)) = messages_stream.next() => {
match &data {
Message::VersionMismatch(VersionMismatchMessage { version, .. }) => {
eprintln!("Version mismatch {} != {}", SRS_VERSION, version);
},
Message::Sync(SyncMessage { clients, .. }) => {
for client in clients.iter() {
client_sink.send(client.clone()).ok();
}
}
Message::RadioUpdate(RadioUpdateMessage { client, .. })=> {
client_sink.send(client.clone()).ok();
}
_ => {},
}
}
Some(data) = voice_handle.recv() => {
frames.push(synthesize(data.as_str(), &voice).await?).await;
}
Some(Ok(data)) = voice_stream.next() => {
// Collect voice packets
let (data, _) = data;
let unit_id = data.unit_id;
let mut t = transmissions.write().await;
if let Entry::Vacant(e) = t.entry(unit_id)
{
e.insert(vec![data]);
} else {
let t1 = t.get_mut(&unit_id).unwrap();
t1.push(data);
}
}
_ = tokio::time::sleep(Duration::from_millis(200)) => {
// Process voice packets
if transmissions.read().await.is_empty() {
continue;
}
let mut transmissions = transmissions.write().await;
for (_, data) in transmissions.iter() {
let mut wav_data = vec![];
let mut unit_sguid = String::new();
for d in data {
unit_sguid = String::from_utf8_lossy(&d.client_sguid).to_string();
if let Some(wav) = &d.wav_audio_part {
wav_data.extend_from_slice(wav);
}
}
let Some(message) = whisper(&stt_url, &wav_data).await else {
continue;
};
tx.send(ReceivedMessage {
unit_sguid,
message
}).ok();
}
transmissions.clear();
}
_ = voice_ping_interval.tick() => {
voice_sink.send((voice_codec::Packet::Ping(sguid), addr)).await?;
}
let Ok(tcp) = TcpStream::connect(addr).await else {
eprintln!("Connection failed: {}", addr);
sleep(Duration::from_secs(5)).await;
continue;
};
let (tcp_stream, tcp_sink) = tcp.into_split();
let mut messages_sink = FramedWrite::new(tcp_sink, MessagesCodec::new());
let mut messages_stream = FramedRead::new(tcp_stream, MessagesCodec::new());
let udp = UdpSocket::bind(SocketAddr::from(([0, 0, 0, 0], 0))).await?;
udp.connect(addr).await?;
let mut voice_ping_interval = tokio::time::interval(Duration::from_secs(15));
let (mut voice_sink, mut voice_stream) = UdpFramed::new(udp, VoiceCodec::new()).split();
let mut packet_id = 1;
if let Err(e) = messages_sink.send(MessageRequest::Sync(SyncMessageRequest {
msg_type: MsgType,
client: client.clone(),
version: SRS_VERSION.to_string(),
})).await {
eprintln!("Srs error: {:?}", e);
sleep(Duration::from_secs(5)).await;
continue;
};
let transmissions = Arc::new(RwLock::new(HashMap::<u32, Vec<VoicePacket>>::new()));
let frames: FrameQueue<Vec<Vec<u8>>> = FrameQueue::new();
loop {
tokio::select! {
Some(data) = frames.next() => {
let start = tokio::time::Instant::now();
for (i, frame) in data.into_iter().enumerate() {
if frame.is_empty() {
continue;
}
let packet = VoicePacket {
audio_part: frame,
wav_audio_part: None,
frequencies: vec![Frequency {
freq: frequency as f64,
modulation: if frequency <= 87_995_000 {
voice_codec::Modulation::Fm
} else {
voice_codec::Modulation::Am
},
encryption: Encryption::None,
}],
unit_id: id.into(),
packet_id,
hop_count: 0,
transmission_sguid: sguid,
client_sguid: sguid,
};
voice_sink.send((packet.into(), addr)).await.ok();
packet_id = packet_id.wrapping_add(1);
let playtime = Duration::from_millis((i as u64 + 1) * 20); // 20m per frame count
let elapsed = start.elapsed();
if playtime > elapsed {
let s = playtime - elapsed;
sleep(s).await;
}
}
}
Some(data) = message_handle.recv() => {
if messages_sink.send(data).await.is_err() {
break;
};
}
Some(data) = messages_stream.next() => {
match data {
Ok(data) => {
match &data {
Message::VersionMismatch(VersionMismatchMessage { version, .. }) => {
eprintln!("Version mismatch {} != {}", SRS_VERSION, version);
},
Message::Sync(SyncMessage { clients, .. }) => {
for client in clients.iter() {
client_sink.send(client.clone()).ok();
}
}
Message::RadioUpdate(RadioUpdateMessage { client, .. })=> {
client_sink.send(client.clone()).ok();
}
_ => {},
}
}
Err(_) => {
break;
}
}
}
Some(data) = voice_handle.recv() => {
frames.push(synthesize(data.as_str(), &voice).await?).await;
}
Some(data) = voice_stream.next() => {
match data {
Ok(data) => {
// Collect voice packets
let (data, _) = data;
let unit_id = data.unit_id;
let mut t = transmissions.write().await;
if let Entry::Vacant(e) = t.entry(unit_id)
{
e.insert(vec![data]);
} else {
let t1 = t.get_mut(&unit_id).unwrap();
t1.push(data);
}
},
Err(_) => {
break;
}
}
}
_ = tokio::time::sleep(Duration::from_millis(200)) => {
// Process voice packets
if transmissions.read().await.is_empty() {
continue;
}
let mut transmissions = transmissions.write().await;
for (_, data) in transmissions.iter() {
let mut wav_data = vec![];
let mut unit_sguid = String::new();
for d in data {
unit_sguid = String::from_utf8_lossy(&d.client_sguid).to_string();
if let Some(wav) = &d.wav_audio_part {
wav_data.extend_from_slice(wav);
}
}
let Some(message) = whisper(&stt_url, &wav_data).await else {
continue;
};
tx.send(ReceivedMessage {
unit_sguid,
message
}).ok();
}
transmissions.clear();
}
_ = voice_ping_interval.tick() => {
if voice_sink.send((voice_codec::Packet::Ping(sguid), addr)).await.is_err() {
break;
};
}
};
}
eprintln!("Disconnected from SRS: {}", addr);
}
#[allow(unreachable_code)]
Ok(())