diff --git a/src/handlers/chat.rs b/src/handlers/chat.rs
index 28b9e20..b91568a 100644
--- a/src/handlers/chat.rs
+++ b/src/handlers/chat.rs
@@ -1,28 +1,28 @@
use axum::{
- extract::{Path, State, Form},
+ extract::{ConnectInfo, Form, Path, State},
response::Html,
};
use chrono::Utc;
-use std::sync::Arc;
+use std::{net::SocketAddr, sync::Arc};
use uuid::Uuid;
use crate::{
- models::{ChatSession, Message, MessageSender, SendMessageForm, NewSessionForm, SessionStatus},
- state::AppState,
+ models::{ChatSession, Message, MessageSender, NewSessionForm, SendMessageForm, SessionStatus},
+ state::{AppState, MESSAGE_LIMIT_PER_WINDOW, SESSION_LIMIT_PER_WINDOW},
};
+const MAX_MESSAGE_CHARS: usize = 1_000;
+const MAX_NAME_CHARS: usize = 40;
+
fn render_message(msg: &Message) -> String {
let (bubble_class, label) = match msg.sender {
- MessageSender::Customer => ("msg-customer", "You"),
- MessageSender::Agent => ("msg-agent", "Support"),
+ MessageSender::User => ("msg-user", "You"),
+ MessageSender::Supporter => ("msg-supporter", "Support"),
MessageSender::System => ("msg-system", ""),
};
if msg.sender == MessageSender::System {
- format!(
- "
{}
",
- html_escape(&msg.content)
- )
+ render_system_notice(&msg.content)
} else {
let time = msg.timestamp.format("%H:%M").to_string();
format!(
@@ -39,11 +39,34 @@ fn render_message(msg: &Message) -> String {
}
}
+fn render_system_notice(content: &str) -> String {
+ format!(
+ "{}
",
+ html_escape(content)
+ )
+}
+
fn html_escape(s: &str) -> String {
s.replace('&', "&")
.replace('<', "<")
.replace('>', ">")
.replace('"', """)
+ .replace('\'', "'")
+}
+
+fn limit_key(kind: &str, addr: SocketAddr) -> String {
+ format!("{}:{}", kind, addr.ip())
+}
+
+fn clean_text(input: &str, max_chars: usize) -> Result {
+ let trimmed = input.trim();
+ if trimmed.is_empty() {
+ return Ok(String::new());
+ }
+ if trimmed.chars().count() > max_chars {
+ return Err(format!("Please keep this under {max_chars} characters."));
+ }
+ Ok(trimmed.to_string())
}
/// GET /messages/:session_id
@@ -52,76 +75,119 @@ pub async fn get_messages(
State(state): State>,
) -> Html {
let Some(session) = state.sessions.get(&session_id) else {
- return Html("Session not found.
".to_string());
+ return Html(render_system_notice(
+ "Session not found. Chat history may have been cleared; please start a new chat.",
+ ));
};
let html: String = session.messages.iter().map(render_message).collect();
Html(html)
}
-/// POST /send/:session_id — customer sends a message
+/// POST /send/:session_id — user sends a message
pub async fn send_message(
+ ConnectInfo(addr): ConnectInfo,
Path(session_id): Path,
State(state): State>,
Form(form): Form,
) -> Html {
- let content = form.content.trim().to_string();
+ if state.is_rate_limited(limit_key("message", addr), MESSAGE_LIMIT_PER_WINDOW) {
+ return Html(render_system_notice(
+ "Too many messages from this network. Please wait until the 3-minute window resets.",
+ ));
+ }
+
+ let content = match clean_text(&form.content, MAX_MESSAGE_CHARS) {
+ Ok(content) => content,
+ Err(message) => return Html(render_system_notice(&message)),
+ };
if content.is_empty() {
return Html(String::new());
}
+
+ let Some(mut session) = state.sessions.get_mut(&session_id) else {
+ return Html(render_system_notice(
+ "Session not found. Chat history may have been cleared; please start a new chat.",
+ ));
+ };
+
let msg = Message {
id: Uuid::new_v4().to_string(),
session_id: session_id.clone(),
- sender: MessageSender::Customer,
+ sender: MessageSender::User,
content,
timestamp: Utc::now(),
};
let rendered = render_message(&msg);
- if let Some(mut session) = state.sessions.get_mut(&session_id) {
- session.messages.push(msg);
- if let Some(tx) = state.notifiers.get(&session_id) {
- let _ = tx.send(session_id.clone());
- }
- let _ = state.agent_notifier.send(session_id.clone());
+ session.messages.push(msg);
+
+ if let Some(tx) = state.notifiers.get(&session_id) {
+ let _ = tx.send(session_id.clone());
}
+ let _ = state.supporter_notifier.send(session_id.clone());
+
Html(rendered)
}
-/// POST /agent/send/:session_id — agent sends a message
-pub async fn agent_send_message(
+/// POST /supporter/send/:session_id — supporter sends a message
+pub async fn supporter_send_message(
+ ConnectInfo(addr): ConnectInfo,
Path(session_id): Path,
State(state): State>,
Form(form): Form,
) -> Html {
- let content = form.content.trim().to_string();
+ if state.is_rate_limited(limit_key("supporter-message", addr), MESSAGE_LIMIT_PER_WINDOW) {
+ return Html(render_system_notice(
+ "Too many messages from this network. Please wait until the 3-minute window resets.",
+ ));
+ }
+
+ let content = match clean_text(&form.content, MAX_MESSAGE_CHARS) {
+ Ok(content) => content,
+ Err(message) => return Html(render_system_notice(&message)),
+ };
if content.is_empty() {
return Html(String::new());
}
+
+ let Some(mut session) = state.sessions.get_mut(&session_id) else {
+ return Html(render_system_notice(
+ "Session not found. Chat history may have been cleared.",
+ ));
+ };
+
let msg = Message {
id: Uuid::new_v4().to_string(),
session_id: session_id.clone(),
- sender: MessageSender::Agent,
+ sender: MessageSender::Supporter,
content,
timestamp: Utc::now(),
};
let rendered = render_message(&msg);
- if let Some(mut session) = state.sessions.get_mut(&session_id) {
- session.messages.push(msg);
- session.status = SessionStatus::Active;
- if let Some(tx) = state.notifiers.get(&session_id) {
- let _ = tx.send(session_id.clone());
- }
+ session.messages.push(msg);
+ session.status = SessionStatus::Active;
+
+ if let Some(tx) = state.notifiers.get(&session_id) {
+ let _ = tx.send(session_id.clone());
}
+ let _ = state.supporter_notifier.send(session_id.clone());
+
Html(rendered)
}
-/// GET /sessions — session list for agent dashboard
-pub async fn get_sessions(
- State(state): State>,
-) -> Html {
- let mut sessions: Vec<_> = state.sessions
+/// GET /sessions — session list for supporter dashboard
+pub async fn get_sessions(State(state): State>) -> Html {
+ let mut sessions: Vec<_> = state
+ .sessions
.iter()
.filter(|s| s.status != SessionStatus::Closed)
- .map(|s| (s.id.clone(), s.customer_name.clone(), s.status.clone(), s.messages.len()))
+ .map(|s| {
+ (
+ s.id.clone(),
+ s.user_name.clone(),
+ s.status.clone(),
+ s.messages.len(),
+ )
+ })
.collect();
sessions.sort_by(|a, b| a.0.cmp(&b.0));
@@ -134,22 +200,24 @@ pub async fn get_sessions(
for (id, name, status, count) in &sessions {
let status_class = match status {
SessionStatus::Waiting => "status-waiting",
- SessionStatus::Active => "status-active",
- SessionStatus::Closed => "status-closed",
+ SessionStatus::Active => "status-active",
+ SessionStatus::Closed => "status-closed",
};
let status_label = match status {
SessionStatus::Waiting => "Waiting",
- SessionStatus::Active => "Active",
- SessionStatus::Closed => "Closed",
+ SessionStatus::Active => "Active",
+ SessionStatus::Closed => "Closed",
};
let short_id = &id[..8];
let escaped_name = html_escape(name);
html.push_str(&format!(
"