diff --git a/Cargo.lock b/Cargo.lock index f37b2fd..f5980f0 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -1734,6 +1734,8 @@ dependencies = [ "tar", "tempfile", "toml", + "unicode-segmentation", + "unicode-width", "ureq", ] diff --git a/Cargo.toml b/Cargo.toml index ea46a69..65ec6a6 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -23,6 +23,8 @@ tar = { version = "0.4.46", default-features = false } toml = "1.1.4" ureq = "3.3.0" tempfile = "3.27.0" +unicode-segmentation = "1.13.3" +unicode-width = "0.2.2" [dev-dependencies] assert_cmd = "2.2.2" diff --git a/src/campaign/campaign_tests.rs b/src/campaign/campaign_tests.rs index a7d31b7..6edb841 100644 --- a/src/campaign/campaign_tests.rs +++ b/src/campaign/campaign_tests.rs @@ -115,6 +115,31 @@ fn the_log_moves_forward_only() { assert!(error.contains("not after the current clock")); } +#[test] +fn a_mark_is_rejected_against_the_older_days_anchor_when_the_newest_day_file_is_empty() { + let (campaign, _dir) = campaign(); + + campaign + .mark(time("#d1-1200"), "first", Some("first")) + .unwrap(); + campaign + .mark(time("#d2-0100"), "second", Some("second")) + .unwrap(); + // A hand-truncated day file: day 2 exists but holds nothing, so the + // real clock is day 1's #d1-1200, not day 2's #d2-0100. + fs::write( + format!("{}/campaign-log/0002.md", _dir.path().display()), + "", + ) + .unwrap(); + + assert_eq!(campaign.current_time().unwrap(), time("#d1-1200")); + let error = campaign + .mark(time("#d1-1200"), "Backdated.", Some("Backdated.")) + .unwrap_err(); + assert!(error.contains("not after the current clock")); +} + #[test] fn the_log_refuses_a_time_before_the_campaign_starts() { let (campaign, _dir) = campaign(); @@ -229,6 +254,24 @@ fn the_transcript_file_has_one_heading_per_section() { assert!(text.contains("## #d1-1200\nNoon arrives.\n")); } +#[test] +fn a_markdown_heading_line_inside_player_input_is_not_mistaken_for_a_new_section() { + let (campaign, _dir) = campaign(); + + campaign + .append_player( + time("#d1-0830"), + "I read the sign.\n## Danger\nWatch your step.", + ) + .unwrap(); + + let sections = campaign.transcript_entries().unwrap(); + assert_eq!(sections.len(), 1); + assert_eq!(sections[0].time, time("#d1-0830")); + assert!(sections[0].body.contains("## Danger")); + assert!(sections[0].body.contains("Watch your step.")); +} + #[test] fn narration_appends_to_the_transcript_alone() { let (campaign, _dir) = campaign(); diff --git a/src/campaign/log.rs b/src/campaign/log.rs index e82def0..87efd7e 100644 --- a/src/campaign/log.rs +++ b/src/campaign/log.rs @@ -30,17 +30,23 @@ impl CampaignLog { &self.dir } - /// The newest `#dX-HHMM` anchor on disk, or `None` when no entry - /// exists yet. + /// The newest `#dX-HHMM` anchor on disk, or `None` when no day file + /// holds one. + /// + /// Walks day files newest first. A day file with no parseable anchor, + /// truncated, emptied, or hand-edited into something else, is not the + /// clock; the search carries on to the next-oldest day rather than + /// reporting no anchor at all when an older day still has one. pub fn last_anchor(&self) -> Result, String> { - let days = days_in(&self.dir); - let Some(&max_day) = days.last() else { - return Ok(None); - }; - let path = day_path(&self.dir, max_day); - let text = fs::read_to_string(&path) - .map_err(|error| format!("{}: cannot read: {error}", path.display()))?; - Ok(parse_lines(&text).last().map(|entry| entry.time)) + for day in days_in(&self.dir).into_iter().rev() { + let path = day_path(&self.dir, day); + let text = fs::read_to_string(&path) + .map_err(|error| format!("{}: cannot read: {error}", path.display()))?; + if let Some(entry) = parse_lines(&text).last() { + return Ok(Some(entry.time)); + } + } + Ok(None) } /// Appends an entry `{time} - {event}` to that day's file, creating it diff --git a/src/campaign/log_tests.rs b/src/campaign/log_tests.rs index 8eb0dc4..7df0080 100644 --- a/src/campaign/log_tests.rs +++ b/src/campaign/log_tests.rs @@ -72,6 +72,40 @@ fn last_anchor_is_the_newest_entry_of_the_highest_day() { ); } +#[test] +fn last_anchor_falls_back_to_an_older_day_when_the_newest_day_file_is_empty() { + let (log, _dir) = log(); + + log.append(GameTime::parse("#d1-0830").unwrap(), "first") + .unwrap(); + // A hand-truncated day file: the day exists but holds nothing. + fs::write(day_path(log.dir(), 2), "").unwrap(); + + assert_eq!( + log.last_anchor().unwrap(), + Some(GameTime::parse("#d1-0830").unwrap()) + ); +} + +#[test] +fn last_anchor_falls_back_to_an_older_day_when_the_newest_day_file_has_no_parseable_lines() { + let (log, _dir) = log(); + + log.append(GameTime::parse("#d1-0830").unwrap(), "first") + .unwrap(); + // A hand-edited day file: text that never matches `#dX-HHMM - text`. + fs::write( + day_path(log.dir(), 2), + "the party presses on\nnothing else worth noting\n", + ) + .unwrap(); + + assert_eq!( + log.last_anchor().unwrap(), + Some(GameTime::parse("#d1-0830").unwrap()) + ); +} + #[test] fn read_lines_returns_all_entries_in_order() { let (log, _dir) = log(); diff --git a/src/campaign/transcript.rs b/src/campaign/transcript.rs index cdd3560..19d218f 100644 --- a/src/campaign/transcript.rs +++ b/src/campaign/transcript.rs @@ -6,7 +6,8 @@ //! paragraphs, and event lines from public or screened marks. Oldest //! entry first. -use std::fs; +use std::fs::{self, OpenOptions}; +use std::io::Write; use std::path::{Path, PathBuf}; use super::clock::GameTime; @@ -37,24 +38,46 @@ impl Transcript { /// change inserts a blank line then the heading. The caller must /// format the block (e.g. `player> …` for the player's lines, or the /// narration as-is). + /// + /// This only ever adds bytes: the existing file is read to decide + /// whether a heading is needed and how much of a gap closes it off + /// cleanly, but nothing already on disk is rewritten, so a crash or a + /// full disk mid-write can only lose the bytes of this one append, + /// never the day's earlier entries. pub fn append(&self, time: GameTime, block: &str) -> Result<(), String> { ensure_dir(&self.dir)?; let path = day_path(&self.dir, time.day); - let (prefix, header) = if path.exists() { - let existing = fs::read_to_string(&path) - .map_err(|error| format!("{}: cannot read: {error}", path.display()))?; - if last_section_time(&existing) == Some(time) { - // Same section; no new heading, no extra separator — - // the file already ends with \n\n. - (existing, String::new()) - } else { - (existing, format!("\n## {time}\n")) - } + let existing = if path.exists() { + Some( + fs::read_to_string(&path) + .map_err(|error| format!("{}: cannot read: {error}", path.display()))?, + ) } else { - (format!("# Day {}\n\n", time.day), format!("## {time}\n")) + None }; - fs::write(&path, format!("{prefix}{header}{block}\n\n")) - .map_err(|error| format!("{}: cannot create: {error}", path.display()))?; + + let mut addition = String::new(); + match existing.as_deref() { + None | Some("") => { + addition.push_str(&format!("# Day {}\n\n## {time}\n", time.day)); + } + Some(text) if last_section_time(text) == Some(time) => { + addition.push_str(gap(text)); + } + Some(text) => { + addition.push_str(gap(text)); + addition.push_str(&format!("## {time}\n")); + } + } + addition.push_str(block); + addition.push_str("\n\n"); + + OpenOptions::new() + .create(true) + .append(true) + .open(&path) + .and_then(|mut file| file.write_all(addition.as_bytes())) + .map_err(|error| format!("{}: cannot append: {error}", path.display()))?; Ok(()) } @@ -77,24 +100,48 @@ fn last_section_time(text: &str) -> Option { parse_sections(text).last().map(|entry| entry.time) } +/// The newlines that close out the end of `text` so the next block starts +/// exactly one blank line below whatever is already on disk. +/// +/// A day file always ends in a blank line, but a hand edit can leave it +/// ending mid-line, with a single trailing newline, or with a blank line +/// already there; each case needs a different number of newlines added to +/// reach the same one-blank-line gap, and an already-ragged tail with +/// extra blank lines is left alone rather than padded further. +fn gap(text: &str) -> &'static str { + match text.len() - text.trim_end_matches('\n').len() { + 0 => "\n\n", + 1 => "\n", + _ => "", + } +} + /// Parses a transcript day-file into its sections: `(time, content)` /// for each `## #dX-HHMM` heading, in order. The `# Day N` title and /// blank lines are ignored. +/// +/// A line only opens a section when it starts with `## ` AND the rest of +/// it parses as a game time; any other `## ` line, a markdown heading +/// written into narration or player input, is ordinary content that +/// belongs to whatever section is already open. fn parse_sections(text: &str) -> Vec { let mut entries = Vec::new(); let mut current: Option = None; let mut body = String::new(); for line in text.lines() { - if line.starts_with("## ") { - if let Some(time) = current.take() { + let heading = line + .strip_prefix("## ") + .and_then(|remainder| GameTime::parse(remainder.trim())); + if let Some(time) = heading { + if let Some(previous) = current.take() { entries.push(LogEntry { - time, + time: previous, body: body.trim().to_string(), }); body.clear(); } - current = GameTime::parse(line.trim_start_matches("## ").trim()); + current = Some(time); } else if current.is_some() { body.push_str(line); body.push('\n'); diff --git a/src/campaign/transcript_tests.rs b/src/campaign/transcript_tests.rs index 33cab52..8e4ed04 100644 --- a/src/campaign/transcript_tests.rs +++ b/src/campaign/transcript_tests.rs @@ -77,6 +77,41 @@ fn read_sections_returns_every_section_in_order() { assert!(sections[1].body.contains("Noon arrives.")); } +#[test] +fn a_markdown_heading_in_a_block_is_not_mistaken_for_a_new_section() { + let (tr, _dir) = t(); + + tr.append( + GameTime::parse("#d1-0830").unwrap(), + "## The Cellar\nDamp stone stairs descend into the dark.", + ) + .unwrap(); + + let sections = tr.read_sections().unwrap(); + assert_eq!(sections.len(), 1); + assert_eq!(sections[0].time.to_string(), "#d1-0830"); + assert!(sections[0].body.contains("## The Cellar")); + assert!(sections[0].body.contains("Damp stone stairs")); +} + +#[test] +fn a_markdown_heading_in_a_block_does_not_confuse_the_next_append() { + let (tr, _dir) = t(); + + tr.append( + GameTime::parse("#d1-0830").unwrap(), + "## The Cellar\nDamp stairs.", + ) + .unwrap(); + tr.append(GameTime::parse("#d1-0830").unwrap(), "You proceed further.") + .unwrap(); + + let sections = tr.read_sections().unwrap(); + assert_eq!(sections.len(), 1); + assert!(sections[0].body.contains("## The Cellar")); + assert!(sections[0].body.contains("You proceed further.")); +} + #[test] fn append_to_a_new_day_starts_a_new_day_file() { let (tr, _dir) = t(); @@ -120,7 +155,7 @@ fn read_sections_fails_when_a_day_file_is_a_directory() { } #[test] -fn append_fails_when_write_fails() { +fn append_fails_when_the_day_file_cannot_be_opened_for_writing() { use std::os::unix::fs::PermissionsExt; let dir = tempfile::TempDir::new().unwrap(); let tr = Transcript::new(dir.path().join("x")); @@ -129,15 +164,74 @@ fn append_fails_when_write_fails() { .unwrap(); let path = day_path(tr.dir(), 1); assert!(path.exists()); - // Chmod the file read-only; the next append will read (succeeds) then - // write (fails). + // Chmod the file read-only; the next append reads it fine but then + // fails to reopen it for the append itself. std::fs::set_permissions(&path, PermissionsExt::from_mode(0o400)).unwrap(); let error = tr .append(GameTime::parse("#d1-0830").unwrap(), "second block") .unwrap_err(); - assert!(error.contains("cannot create")); // fs::write error + assert!(error.contains("cannot append")); +} + +#[test] +fn append_normalizes_a_hand_removed_trailing_blank_line() { + let (tr, _dir) = t(); + let path = day_path(tr.dir(), 1); + day_file::ensure_dir(tr.dir()).unwrap(); + // A hand-edited file: the blank line the transcript always ends in + // was trimmed away, leaving a single trailing newline. + fs::write(&path, "# Day 1\n\n## #d1-0830\nplayer> I wake up.\n").unwrap(); + + tr.append( + GameTime::parse("#d1-0830").unwrap(), + "You find dusty cobwebs.", + ) + .unwrap(); + + let text = fs::read_to_string(&path).unwrap(); + assert_eq!( + text, + "# Day 1\n\n## #d1-0830\nplayer> I wake up.\n\nYou find dusty cobwebs.\n\n" + ); +} + +#[test] +fn append_normalizes_a_hand_edited_file_with_no_trailing_newline_at_all() { + let (tr, _dir) = t(); + let path = day_path(tr.dir(), 1); + day_file::ensure_dir(tr.dir()).unwrap(); + // A hand-edited file: both trailing newlines are gone, so the file + // ends mid-line. + fs::write(&path, "# Day 1\n\n## #d1-0830\nplayer> I wake up.").unwrap(); + + tr.append( + GameTime::parse("#d1-0830").unwrap(), + "You find dusty cobwebs.", + ) + .unwrap(); + + let text = fs::read_to_string(&path).unwrap(); + assert_eq!( + text, + "# Day 1\n\n## #d1-0830\nplayer> I wake up.\n\nYou find dusty cobwebs.\n\n" + ); +} + +#[test] +fn append_does_not_rewrite_bytes_already_on_disk() { + let (tr, _dir) = t(); + tr.append(GameTime::parse("#d1-0830").unwrap(), "first block.") + .unwrap(); + let path = day_path(tr.dir(), 1); + let before = fs::read_to_string(&path).unwrap(); + + tr.append(GameTime::parse("#d1-1200").unwrap(), "second block.") + .unwrap(); + + let after = fs::read_to_string(&path).unwrap(); + assert!(after.starts_with(&before)); } #[test] diff --git a/src/chat/client.rs b/src/chat/client.rs index d5dbd66..5dd71b1 100644 --- a/src/chat/client.rs +++ b/src/chat/client.rs @@ -1,6 +1,7 @@ use std::collections::VecDeque; use std::fmt; use std::io; +use std::time::Duration; use crate::chat::sse::SseReader; use crate::chat::wire::{ @@ -10,6 +11,53 @@ use crate::chat::wire::{ /// The maximum length of a status error's body, in characters. const STATUS_ERROR_BODY_LIMIT: usize = 2_000; +/// How many times the client sends one request before it gives up. +const MAX_ATTEMPTS: u32 = 3; + +/// The time limits on one chat completion request. +/// +/// `stream` uses the defaults. The tests build shorter ones so a stalled +/// server or a backoff does not hold the suite. +#[derive(Clone, Copy)] +struct Limits { + /// How long the socket and the TLS handshake may take. + connect: Duration, + /// How long the client waits for the response headers after it sends + /// the request. + response: Duration, + /// How long the whole streamed body may take. + body: Duration, + /// How long to wait before the second attempt. Each further attempt + /// doubles it. + backoff: Duration, +} + +impl Default for Limits { + fn default() -> Self { + Self { + connect: Duration::from_secs(10), + response: Duration::from_secs(60), + // ureq measures this over the whole body, and the body of a + // streamed turn is the whole turn. So this is a bound on a + // dead connection, not a budget for the model: a healthy turn + // finishes well inside 10 minutes, and a provider that stops + // writing mid-stream frees the worker thread at 10 minutes + // rather than never. + body: Duration::from_secs(600), + backoff: Duration::from_millis(500), + } + } +} + +/// Reports whether a status is worth sending the request again. +/// +/// 429 means the provider is rate limiting, and a 5xx is a failure on its +/// side that the next attempt may not hit. Every other status describes +/// the request itself, so a second attempt gives the same answer. +fn worth_retrying(status: u16) -> bool { + status == 429 || (500..600).contains(&status) +} + /// A client for one chat completion endpoint. pub struct Client { pub api_base: String, @@ -23,10 +71,26 @@ impl Client { /// Sends `messages` with `self.model` to `{api_base}/chat/completions`. /// A trailing slash on `api_base` is tolerated. An empty `tools` omits /// the field from the request rather than sending an empty array. + /// + /// A rate limit or a server error costs one of `MAX_ATTEMPTS` tries, + /// with a doubling wait between them. Retrying here is safe because no + /// byte has reached the caller yet: the caller sees one request, and a + /// failure leaves it with nothing to undo. Once the body starts, the + /// client never sends the request again. pub fn stream( &self, messages: &[Message], tools: &[serde_json::Value], + ) -> Result { + self.stream_within(messages, tools, Limits::default()) + } + + /// Runs `stream` under `limits`. + fn stream_within( + &self, + messages: &[Message], + tools: &[serde_json::Value], + limits: Limits, ) -> Result { let request = ChatRequest { model: self.model.clone(), @@ -38,29 +102,42 @@ impl Client { let config = ureq::Agent::config_builder() .http_status_as_error(false) + .timeout_connect(Some(limits.connect)) + .timeout_recv_response(Some(limits.response)) + .timeout_recv_body(Some(limits.body)) .build(); let agent: ureq::Agent = config.into(); let base = self.api_base.trim_end_matches('/'); let url = format!("{base}/chat/completions"); - let mut response = agent - .post(&url) - .header("Authorization", format!("Bearer {}", self.api_key)) - .header("Content-Type", "application/json") - .send(body.as_slice()) - .map_err(ChatError::Transport)?; - - let status = response.status(); - if !status.is_success() { + + let mut attempts_left = MAX_ATTEMPTS - 1; + let mut backoff = limits.backoff; + loop { + let mut response = agent + .post(&url) + .header("Authorization", format!("Bearer {}", self.api_key)) + .header("Content-Type", "application/json") + .send(body.as_slice()) + .map_err(ChatError::Transport)?; + + let status = response.status(); + if status.is_success() { + return Ok(ChatStream::new(response.into_body().into_reader())); + } + let body_text = response.body_mut().read_to_string().unwrap_or_default(); let truncated: String = body_text.chars().take(STATUS_ERROR_BODY_LIMIT).collect(); - return Err(ChatError::Status { - status: status.as_u16(), - body: truncated, - }); + if attempts_left == 0 || !worth_retrying(status.as_u16()) { + return Err(ChatError::Status { + status: status.as_u16(), + body: truncated, + }); + } + attempts_left -= 1; + std::thread::sleep(backoff); + backoff *= 2; } - - Ok(ChatStream::new(response.into_body().into_reader())) } } @@ -77,19 +154,34 @@ pub enum StreamItem { /// each tool call as soon as its fragments are complete. /// /// The stream ends when the server sends `[DONE]` or closes the connection. -/// The OpenAI streaming contract sends a tool call's fragments in -/// non-decreasing index order, so a call is complete, and ready to yield, -/// the moment a fragment for a higher index arrives; the call at the -/// highest index seen completes when the stream ends. +/// +/// OpenAI-compatible providers number tool calls in ways that disagree. +/// Some give every call the index 0, and some send the fragments of two +/// calls one after the other, 0, 1, 0, 1. So the index is a hint, not a +/// key: a fragment that carries an id joins the call with that id, and a +/// fragment without one joins the newest open call at its index. An id no +/// open call has starts a new call, even at an index already in use. +/// +/// A call leaves the stream when a later call starts and its arguments +/// parse as one whole JSON value, so the caller can run a tool while the +/// model still writes the next call. Every call still open when the stream +/// ends leaves then. pub struct ChatStream { events: SseReader>, - /// The call at the highest index seen so far, still being assembled. - current_call: Option<(usize, PartialToolCall)>, + /// The calls still being assembled, oldest first. + open: VecDeque, /// Items ready to yield before the next chunk is read: one chunk can /// both complete a call and carry text. queue: VecDeque, } +/// A tool call being assembled from fragments, with the index the provider +/// gave it. +struct OpenToolCall { + index: usize, + partial: PartialToolCall, +} + /// A tool call being assembled from fragments. #[derive(Default)] struct PartialToolCall { @@ -99,6 +191,15 @@ struct PartialToolCall { } impl PartialToolCall { + /// Reports whether the arguments seen so far are one whole JSON value. + /// + /// A provider that splits arguments sends a prefix of a JSON value, + /// and a prefix never parses. So a parse that succeeds means every + /// argument fragment has arrived. + fn arguments_are_whole(&self) -> bool { + serde_json::from_str::(&self.arguments).is_ok() + } + fn into_tool_call(self) -> ToolCall { ToolCall { id: self.id, @@ -115,44 +216,71 @@ impl ChatStream { fn new(reader: ureq::BodyReader<'static>) -> Self { Self { events: SseReader::new(reader), - current_call: None, + open: VecDeque::new(), queue: VecDeque::new(), } } - /// Merges one chunk's tool-call fragments into the call being - /// assembled, queuing the call at the previous index as soon as a - /// fragment names a higher one. + /// Merges one chunk's tool-call fragments into the calls being + /// assembled, queuing finished calls whenever a new call starts. fn absorb_tool_call_fragments(&mut self, fragments: Vec) { for fragment in fragments { - if self - .current_call - .as_ref() - .is_some_and(|(index, _)| *index != fragment.index) - { - let (_, partial) = self - .current_call - .take() - .expect("current_call is Some, just checked above"); - self.queue - .push_back(StreamItem::Call(partial.into_tool_call())); + match self.position_of(&fragment) { + Some(position) => self.merge_into(position, fragment), + None => { + self.queue_finished_calls(); + self.open.push_back(OpenToolCall { + index: fragment.index, + partial: PartialToolCall::default(), + }); + let newest = self.open.len() - 1; + self.merge_into(newest, fragment); + } } + } + } - let (_, partial) = self - .current_call - .get_or_insert_with(|| (fragment.index, PartialToolCall::default())); - if let Some(id) = fragment.id { - partial.id = id; - } - let Some(function) = fragment.function else { - continue; - }; - if let Some(name) = function.name { - partial.name = name; - } - if let Some(arguments) = function.arguments { - partial.arguments.push_str(&arguments); - } + /// Finds the open call `fragment` belongs to, or `None` when it starts + /// a new one. + fn position_of(&self, fragment: &ToolCallDelta) -> Option { + match &fragment.id { + Some(id) => self.open.iter().position(|open| open.partial.id == *id), + None => self + .open + .iter() + .rposition(|open| open.index == fragment.index), + } + } + + /// Copies whatever `fragment` carries into the open call at `position`. + fn merge_into(&mut self, position: usize, fragment: ToolCallDelta) { + let partial = &mut self.open[position].partial; + if let Some(id) = fragment.id { + partial.id = id; + } + let Some(function) = fragment.function else { + return; + }; + if let Some(name) = function.name { + partial.name = name; + } + if let Some(arguments) = function.arguments { + partial.arguments.push_str(&arguments); + } + } + + /// Queues each finished call from the front of `open`, stopping at the + /// first one still waiting on fragments. Stopping there keeps the + /// calls in the order the provider started them. + fn queue_finished_calls(&mut self) { + while self + .open + .front() + .is_some_and(|open| open.partial.arguments_are_whole()) + { + let finished = self.open.pop_front().expect("front is Some, just checked"); + self.queue + .push_back(StreamItem::Call(finished.partial.into_tool_call())); } } } @@ -169,9 +297,9 @@ impl Iterator for ChatStream { let payload = match self.events.next() { None => { return self - .current_call - .take() - .map(|(_, partial)| Ok(StreamItem::Call(partial.into_tool_call()))); + .open + .pop_front() + .map(|open| Ok(StreamItem::Call(open.partial.into_tool_call()))); } Some(Ok(payload)) => payload, Some(Err(error)) => return Some(Err(ChatError::Io(error))), @@ -255,18 +383,53 @@ mod tests { /// `response` must be a complete HTTP/1.1 response, including status /// line and headers. fn fake_server(response: Vec) -> (String, Receiver, JoinHandle<()>) { + fake_server_serving(vec![response]) + } + + /// Serves one response from `responses` per connection, in order, and + /// sends every request it read back through the returned channel. + fn fake_server_serving( + responses: Vec>, + ) -> (String, Receiver, JoinHandle<()>) { let listener = TcpListener::bind("127.0.0.1:0").unwrap(); let addr = listener.local_addr().unwrap(); let (sender, receiver) = mpsc::channel(); let handle = std::thread::spawn(move || { - let (mut stream, _) = listener.accept().unwrap(); - let captured = read_request(&stream); - stream.write_all(&response).unwrap(); - sender.send(captured).unwrap(); + for response in responses { + let (mut stream, _) = listener.accept().unwrap(); + let captured = read_request(&stream); + stream.write_all(&response).unwrap(); + sender.send(captured).unwrap(); + } }); (format!("http://{addr}"), receiver, handle) } + /// Accepts one connection, reads the request, writes `sent`, and then + /// holds the connection open for `hold` without writing anything more. + fn stalling_server(sent: Vec, hold: Duration) -> (String, JoinHandle<()>) { + let listener = TcpListener::bind("127.0.0.1:0").unwrap(); + let addr = listener.local_addr().unwrap(); + let handle = std::thread::spawn(move || { + let (mut stream, _) = listener.accept().unwrap(); + read_request(&stream); + stream.write_all(&sent).unwrap(); + std::thread::sleep(hold); + }); + (format!("http://{addr}"), handle) + } + + /// Limits short enough that a stalled server or a backoff costs the + /// suite a fraction of a second. + fn quick_limits() -> Limits { + Limits { + connect: Duration::from_secs(1), + response: Duration::from_millis(200), + body: Duration::from_millis(200), + backoff: Duration::from_millis(1), + } + } + /// Reads one HTTP/1.1 request's method, path, headers, and body from /// `stream`. fn read_request(stream: &std::net::TcpStream) -> CapturedRequest { @@ -341,6 +504,26 @@ mod tests { "http://127.0.0.1:0".to_string() } + /// The `StreamItem::Call` a finished tool call arrives as. + fn call_item(id: &str, name: &str, arguments: &str) -> StreamItem { + StreamItem::Call(ToolCall { + id: id.to_string(), + kind: "function".to_string(), + function: ToolCallFunction { + name: name.to_string(), + arguments: arguments.to_string(), + }, + }) + } + + /// One SSE event carrying `fragments` as a chunk's `tool_calls`. + fn tool_call_event(fragments: serde_json::Value) -> String { + let chunk = serde_json::json!({ + "choices": [{"delta": {"tool_calls": fragments}, "finish_reason": null}], + }); + format!("data: {chunk}\n\n") + } + fn a_message() -> Vec { vec![Message { role: crate::chat::wire::Role::User, @@ -415,19 +598,98 @@ mod tests { #[test] fn a_non_success_status_becomes_a_status_error_with_the_body() { let (url, _requests, server) = fake_server(http_response( - "500 Internal Server Error", + "400 Bad Request", "text/plain", - "server exploded", + "unknown model", )); let client = client_for(url); let error = client.stream(&a_message(), &[]).err().unwrap(); - assert!(matches!(error, ChatError::Status { status: 500, .. })); + assert!(matches!(error, ChatError::Status { status: 400, .. })); assert_eq!( error.to_string(), - "the chat server answered with status 500: server exploded" + "the chat server answered with status 400: unknown model" + ); + server.join().unwrap(); + } + + #[test] + fn a_rate_limited_request_goes_out_again_and_the_next_answer_streams() { + let rate_limited = || http_response("429 Too Many Requests", "text/plain", "slow down"); + let body = "data: {\"choices\":[{\"delta\":{\"content\":\"Hi\"},\"finish_reason\":null}]}\n\n\ + data: [DONE]\n\n"; + let (url, requests, server) = + fake_server_serving(vec![rate_limited(), rate_limited(), sse_response(body)]); + let client = client_for(url); + + let items: Vec = client + .stream_within(&a_message(), &[], quick_limits()) + .unwrap() + .map(|item| item.unwrap()) + .collect(); + + assert_eq!(items, vec![StreamItem::Text("Hi".to_string())]); + server.join().unwrap(); + assert_eq!(requests.iter().count(), 3); + } + + #[test] + fn a_server_error_on_every_attempt_gives_up_with_the_status() { + let failing = || http_response("503 Service Unavailable", "text/plain", "no capacity"); + let (url, requests, server) = fake_server_serving(vec![failing(), failing(), failing()]); + let client = client_for(url); + + let error = client + .stream_within(&a_message(), &[], quick_limits()) + .err() + .unwrap(); + + assert!(matches!(error, ChatError::Status { status: 503, .. })); + server.join().unwrap(); + assert_eq!(requests.iter().count(), MAX_ATTEMPTS as usize); + } + + #[test] + fn a_server_that_takes_the_request_and_never_answers_fails_the_request() { + let (url, server) = stalling_server(Vec::new(), Duration::from_millis(400)); + let client = client_for(url); + + let error = client + .stream_within(&a_message(), &[], quick_limits()) + .err() + .unwrap(); + + assert!(matches!(error, ChatError::Transport(_))); + server.join().unwrap(); + } + + #[test] + fn a_body_that_stops_arriving_ends_the_stream_with_an_io_error() { + let sse_body = + "data: {\"choices\":[{\"delta\":{\"content\":\"Hi\"},\"finish_reason\":null}]}\n\n"; + let head = format!( + "HTTP/1.1 200 OK\r\n\ + Content-Type: text/event-stream\r\n\ + Content-Length: {}\r\n\ + Connection: close\r\n\ + \r\n\ + {sse_body}", + sse_body.len() + 100, + ) + .into_bytes(); + let (url, server) = stalling_server(head, Duration::from_millis(400)); + let client = client_for(url); + + let mut stream = client + .stream_within(&a_message(), &[], quick_limits()) + .unwrap(); + + assert_eq!( + stream.next().unwrap().unwrap(), + StreamItem::Text("Hi".to_string()) ); + assert!(matches!(stream.next(), Some(Err(ChatError::Io(_))))); server.join().unwrap(); } @@ -576,6 +838,139 @@ mod tests { server.join().unwrap(); } + #[test] + fn two_calls_reported_at_the_same_index_stay_separate() { + let body = format!( + "{}{}data: [DONE]\n\n", + tool_call_event(serde_json::json!([{ + "index": 0, + "id": "call_1", + "function": {"name": "roll", "arguments": "{\"notation\":\"1d20\"}"}, + }])), + tool_call_event(serde_json::json!([{ + "index": 0, + "id": "call_2", + "function": {"name": "roll", "arguments": "{\"notation\":\"1d6\"}"}, + }])), + ); + let (url, _requests, server) = fake_server(sse_response(&body)); + let client = client_for(url); + + let items: Vec = client + .stream(&a_message(), &[]) + .unwrap() + .map(|item| item.unwrap()) + .collect(); + + assert_eq!( + items, + vec![ + call_item("call_1", "roll", "{\"notation\":\"1d20\"}"), + call_item("call_2", "roll", "{\"notation\":\"1d6\"}"), + ] + ); + server.join().unwrap(); + } + + #[test] + fn interleaved_fragments_for_two_indexes_assemble_into_two_whole_calls() { + let body = format!( + "{}{}{}{}data: [DONE]\n\n", + tool_call_event(serde_json::json!([{ + "index": 0, + "id": "call_1", + "function": {"name": "roll", "arguments": "{\"notation\":"}, + }])), + tool_call_event(serde_json::json!([{ + "index": 1, + "id": "call_2", + "function": {"name": "roll", "arguments": "{\"notation\":"}, + }])), + tool_call_event(serde_json::json!([{ + "index": 0, + "function": {"arguments": "\"1d20\"}"}, + }])), + tool_call_event(serde_json::json!([{ + "index": 1, + "function": {"arguments": "\"1d6\"}"}, + }])), + ); + let (url, _requests, server) = fake_server(sse_response(&body)); + let client = client_for(url); + + let items: Vec = client + .stream(&a_message(), &[]) + .unwrap() + .map(|item| item.unwrap()) + .collect(); + + assert_eq!( + items, + vec![ + call_item("call_1", "roll", "{\"notation\":\"1d20\"}"), + call_item("call_2", "roll", "{\"notation\":\"1d6\"}"), + ] + ); + server.join().unwrap(); + } + + #[test] + fn a_whole_call_arriving_in_one_fragment_yields_one_call() { + let body = format!( + "{}data: [DONE]\n\n", + tool_call_event(serde_json::json!([{ + "index": 0, + "id": "call_1", + "function": {"name": "roll", "arguments": "{\"notation\":\"1d20\"}"}, + }])), + ); + let (url, _requests, server) = fake_server(sse_response(&body)); + let client = client_for(url); + + let items: Vec = client + .stream(&a_message(), &[]) + .unwrap() + .map(|item| item.unwrap()) + .collect(); + + assert_eq!( + items, + vec![call_item("call_1", "roll", "{\"notation\":\"1d20\"}")] + ); + server.join().unwrap(); + } + + #[test] + fn a_fragment_that_repeats_the_id_continues_the_same_call() { + let body = format!( + "{}{}data: [DONE]\n\n", + tool_call_event(serde_json::json!([{ + "index": 0, + "id": "call_1", + "function": {"name": "roll", "arguments": "{\"notation\":"}, + }])), + tool_call_event(serde_json::json!([{ + "index": 0, + "id": "call_1", + "function": {"arguments": "\"1d20\"}"}, + }])), + ); + let (url, _requests, server) = fake_server(sse_response(&body)); + let client = client_for(url); + + let items: Vec = client + .stream(&a_message(), &[]) + .unwrap() + .map(|item| item.unwrap()) + .collect(); + + assert_eq!( + items, + vec![call_item("call_1", "roll", "{\"notation\":\"1d20\"}")] + ); + server.join().unwrap(); + } + #[test] fn text_and_completed_calls_interleave_in_stream_order() { let body = "data: {\"choices\":[{\"delta\":{\"content\":\"Hello\"},\"finish_reason\":null}]}\n\n\ diff --git a/src/chat/sse.rs b/src/chat/sse.rs index 69e2d76..07a6413 100644 --- a/src/chat/sse.rs +++ b/src/chat/sse.rs @@ -9,6 +9,10 @@ use std::io::{self, BufRead, BufReader, Read}; /// are read and discarded. pub struct SseReader { lines: BufReader, + /// Set once `[DONE]` or EOF ends the stream, so no later call reads + /// the connection again. On a connection the client reuses, the bytes + /// after the end of one response belong to the next one. + finished: bool, } impl SseReader { @@ -16,6 +20,7 @@ impl SseReader { pub fn new(reader: R) -> Self { Self { lines: BufReader::new(reader), + finished: false, } } } @@ -24,7 +29,14 @@ impl Iterator for SseReader { type Item = io::Result; fn next(&mut self) -> Option { - next_event(&mut self.lines) + if self.finished { + return None; + } + let event = next_event(&mut self.lines); + if event.is_none() { + self.finished = true; + } + event } } @@ -166,6 +178,47 @@ mod tests { assert_eq!(result, vec!["hello".to_string()]); } + #[test] + fn every_call_after_done_returns_none() { + let mut reader = SseReader::new("data: hi\n\ndata: [DONE]\n\ndata: after\n\n".as_bytes()); + + assert_eq!(reader.next().unwrap().unwrap(), "hi"); + assert!(reader.next().is_none()); + assert!(reader.next().is_none()); + } + + /// A reader that hands back one chunk per `read`, where an empty chunk + /// reports EOF. A reused connection behaves this way: the response ends, + /// and the bytes of the next one follow on the same socket. + struct ChunksWithAnEofInTheMiddle { + chunks: std::collections::VecDeque<&'static [u8]>, + } + + impl Read for ChunksWithAnEofInTheMiddle { + fn read(&mut self, buf: &mut [u8]) -> io::Result { + let chunk = self.chunks.pop_front().unwrap_or(b""); + buf[..chunk.len()].copy_from_slice(chunk); + Ok(chunk.len()) + } + } + + #[test] + fn every_call_after_eof_returns_none_even_when_more_bytes_follow() { + let mut reader = SseReader::new(ChunksWithAnEofInTheMiddle { + chunks: [ + b"data: hi\n\n".as_slice(), + b"", + b"data: after\n\n".as_slice(), + ] + .into_iter() + .collect(), + }); + + assert_eq!(reader.next().unwrap().unwrap(), "hi"); + assert!(reader.next().is_none()); + assert!(reader.next().is_none()); + } + struct OneByteAtATime { remaining: std::collections::VecDeque, } diff --git a/src/dm/dm_tests.rs b/src/dm/dm_tests.rs index 73ed1e1..a52c71d 100644 --- a/src/dm/dm_tests.rs +++ b/src/dm/dm_tests.rs @@ -201,6 +201,34 @@ fn a_turn_with_a_campaign_logs_its_narration_to_the_transcript() { assert!(!world.path().join("campaign-log/0001.md").exists()); } +#[test] +fn a_pre_turn_campaign_write_failure_ends_the_turn_as_an_error_before_any_request() { + use std::os::unix::fs::PermissionsExt; + let world = TempDir::new().unwrap(); + let campaign = Campaign::open(world.path()).unwrap(); + // The transcript directory cannot be written to, so recording the + // player's line fails before the turn ever reaches the network. + std::fs::set_permissions( + world.path().join("transcript"), + PermissionsExt::from_mode(0o500), + ) + .unwrap(); + let mut dm = Dm::new( + Config { + api_base: "http://127.0.0.1:0".to_string(), + api_key: "sk-test".to_string(), + model: "gpt-4o-mini".to_string(), + }, + Arc::new(fixtures::mount(&[])), + &[], + Some(campaign), + ); + + let error = dm.turn("I sleep.", &mut ignore_event).unwrap_err(); + + assert!(error.to_string().contains("cannot append")); +} + #[test] fn a_second_turn_sends_the_first_turns_messages_in_the_history() { let first_reply = "data: {\"choices\":[{\"delta\":{\"content\":\"You see a door.\"},\"finish_reason\":null}]}\n\n\ @@ -230,7 +258,7 @@ fn a_second_turn_sends_the_first_turns_messages_in_the_history() { #[test] fn a_failed_turn_leaves_the_history_unchanged() { let (url, requests, server) = fake_server(vec![ - http_response("500 Internal Server Error", "text/plain", "server exploded"), + http_response("400 Bad Request", "text/plain", "unknown model"), sse_response("data: [DONE]\n\n"), ]); let mut dm = dm_for(url); diff --git a/src/dm/dm_tool_round_tests.rs b/src/dm/dm_tool_round_tests.rs index 4fb54c5..f5a67bf 100644 --- a/src/dm/dm_tool_round_tests.rs +++ b/src/dm/dm_tool_round_tests.rs @@ -6,11 +6,13 @@ use super::tests::{CapturedRequest, fake_server, ignore_event, sent_messages, sse_response}; use super::*; +use crate::campaign::GameTime; use crate::context::ContextStack; use crate::knowledge::fixtures; use rand::SeedableRng; use rand::rngs::StdRng; use serde_json::json; +use tempfile::TempDir; fn empty_context() -> ContextStack { let mount = fixtures::mount(&[]); @@ -151,6 +153,66 @@ fn a_tool_round_then_a_reply_the_second_request_carries_the_tool_result() { assert!(messages[3]["content"].as_str().unwrap().contains("total")); } +#[test] +fn a_post_turn_transcript_write_failure_appends_a_warning_instead_of_failing_the_turn() { + use std::os::unix::fs::PermissionsExt; + let mark_args = json!({ + "time": "#d2-0100", + "event": "Morning breaks", + "visibility": "secret", + }) + .to_string(); + let (url, _requests, server) = fake_server(vec![ + round_response("", &[call("call_1", "mark", &mark_args)]), + reply_response("You wake up."), + ]); + let world = TempDir::new().unwrap(); + let campaign = Campaign::open(world.path()).unwrap(); + // Seed day 1 so the turn's own player line only ever appends to a + // day file that already exists. + campaign + .append_narration(GameTime::parse("#d1-0000").unwrap(), "(seed)") + .unwrap(); + // Nothing new can be created in the transcript directory, so day 2's + // file, which the end-of-turn narration would open for the first + // time after the mid-turn mark advances the clock, cannot be + // written. + std::fs::set_permissions( + world.path().join("transcript"), + PermissionsExt::from_mode(0o500), + ) + .unwrap(); + let mut dm = Dm::new( + Config { + api_base: url, + api_key: "sk-test".to_string(), + model: "gpt-4o-mini".to_string(), + }, + Arc::new(fixtures::mount(&[])), + &[], + Some(campaign), + ); + let mut deltas = Vec::new(); + + let turn = dm + .turn("I sleep.", &mut |event| { + if let TurnDelta::Text(text) = event { + deltas.push(text); + } + ControlFlow::Continue(()) + }) + .unwrap(); + + server.join().unwrap(); + let Turn::Reply(reply) = turn else { + panic!("expected Turn::Reply, got {turn:?}"); + }; + assert!(reply.contains("You wake up.")); + assert!(reply.contains("cannot append")); + let warning = deltas.last().expect("a warning delta was streamed"); + assert!(warning.contains("cannot append")); +} + #[test] fn a_completed_multi_round_turns_history_carries_every_round_message_in_order() { let (url, requests, server) = fake_server(vec![ diff --git a/src/dm/mod.rs b/src/dm/mod.rs index ef9c95e..aa14ee3 100644 --- a/src/dm/mod.rs +++ b/src/dm/mod.rs @@ -1,6 +1,7 @@ //! The dungeon master: the system prompt, the message history, and the //! round loop that runs one player turn. +use std::fmt; use std::ops::ControlFlow; use std::path::PathBuf; use std::sync::Arc; @@ -53,13 +54,43 @@ const MAX_ROUNDS: usize = 24; /// What one turn produced. #[derive(Debug, Clone, PartialEq, Eq)] pub enum Turn { - /// The reply completed and joined the history. + /// The reply completed and joined the history. If the campaign could + /// not record the turn's narration after it had already streamed to + /// the player, a warning naming the problem is appended to the reply + /// here; the history keeps the DM's own narration without it. Reply(String), /// The caller stopped the stream partway through. The history is /// unchanged, and the reply collected so far is returned for display. Cancelled(String), } +/// Why a turn ended without a reply: the chat request failed, or the +/// player's line could not be recorded before anything streamed. +#[derive(Debug)] +pub enum TurnError { + /// Sending the request or reading the streamed response failed. + Chat(ChatError), + /// The campaign could not record the turn. + Campaign(String), +} + +impl fmt::Display for TurnError { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + match self { + TurnError::Chat(error) => write!(f, "{error}"), + TurnError::Campaign(message) => write!(f, "{message}"), + } + } +} + +impl std::error::Error for TurnError {} + +impl From for TurnError { + fn from(error: ChatError) -> Self { + TurnError::Chat(error) + } +} + /// One piece of a turn as it runs, passed to `Dm::turn`'s callback. #[derive(Debug, Clone, PartialEq)] pub enum TurnDelta { @@ -180,11 +211,18 @@ impl Dm { /// On `Break`, the history is unchanged and the result is /// `Turn::Cancelled` with the current round's narration collected so /// far. On failure, the history is unchanged. + /// + /// Recording the player's line to the campaign happens before any + /// request is sent, so a failure there ends the turn as an error with + /// nothing shown to the player yet. Recording the narration happens + /// after the reply has already streamed, so a failure there cannot + /// undo what the player saw; it turns into a warning appended to the + /// reply and streamed through `on_delta` like any other text. pub fn turn( &mut self, input: &str, on_delta: &mut dyn FnMut(TurnDelta) -> ControlFlow<()>, - ) -> Result { + ) -> Result { // The system prompt is built fresh: if the mount has changed // since the last turn, the next request carries the latest // fragments and context entries. @@ -213,8 +251,12 @@ impl Dm { if let (Some(campaign), Some(time)) = (&self.campaign, start_time) { // The player's line goes in at the start-of-turn clock so it // appears in chronological order before any mark that - // advances time mid-turn. - let _ = campaign.append_player(time, input); + // advances time mid-turn. Nothing has streamed yet, so a + // failure here ends the turn outright rather than carrying on + // with a campaign record that is already missing a line. + campaign + .append_player(time, input) + .map_err(TurnError::Campaign)?; } loop { @@ -285,8 +327,19 @@ impl Dm { // the turn, after any marks that advanced it, so the // transcript stays oldest-to-newest. let now = campaign.current_time().ok().or(start_time); - if let (Some(time), true) = (now, !turn_narration.trim().is_empty()) { - let _ = campaign.append_narration(time, &turn_narration); + // The narration has already streamed to the player by + // this point, so a write failure here cannot fail the + // turn outright; it becomes a warning appended to the + // reply and streamed the same way the narration was, + // so the player sees it instead of the record silently + // going missing. + if let (Some(time), true) = (now, !turn_narration.trim().is_empty()) + && let Err(error) = campaign.append_narration(time, &turn_narration) + { + let warning = + format!("\n\n(warning: the transcript was not saved: {error})"); + narration.push_str(&warning); + let _ = on_delta(TurnDelta::Text(warning)); } } return Ok(Turn::Reply(narration)); diff --git a/src/markdown/block.rs b/src/markdown/block.rs index e45829f..63cc1ed 100644 --- a/src/markdown/block.rs +++ b/src/markdown/block.rs @@ -15,6 +15,8 @@ use pulldown_cmark::{Event, HeadingLevel, Tag, TagEnd}; use ratatui::style::Style; use ratatui::text::{Line, Span}; +use unicode_segmentation::UnicodeSegmentation; +use unicode_width::UnicodeWidthStr; use super::{Events, inline, join, style, table}; use crate::wrap::wrap_spans; @@ -251,8 +253,11 @@ fn code_block(events: &mut Events, end: TagEnd, width: usize) -> Vec Vec> { let mut source = String::new(); while let Some(event) = events.next() { @@ -265,7 +270,7 @@ fn html_block(events: &mut Events, end: TagEnd, width: usize, base: Style) -> Ve } source_lines(&source) .iter() - .flat_map(|line| hard_wrap(line, width)) + .flat_map(|line| hard_wrap(line, width.max(1))) .map(|row| Line::from(Span::styled(row, base))) .collect() } @@ -281,18 +286,26 @@ fn source_lines(source: &str) -> Vec<&str> { .collect() } -/// Cuts `line` into pieces of at most `width` characters, verbatim, with -/// no whitespace collapsed. A line that fits comes back whole; an empty -/// line comes back as itself, one empty row. +/// Cuts `line` into pieces of at most `width` columns, verbatim, with no +/// whitespace collapsed. A cut lands between two grapheme clusters, so a +/// letter keeps its accent and one emoji stays in one piece. A line that +/// fits comes back whole; an empty line comes back as itself, one empty +/// row. A cluster wider than `width`, which a `width` of zero makes of +/// every cluster, is a row of its own. fn hard_wrap(line: &str, width: usize) -> Vec { - let characters: Vec = line.chars().collect(); - if characters.is_empty() { + if line.is_empty() { return vec![String::new()]; } - characters - .chunks(width) - .map(|chunk| chunk.iter().collect()) - .collect() + let mut rows = Vec::new(); + let mut row = String::new(); + for cluster in line.graphemes(true) { + if !row.is_empty() && row.width() + cluster.width() > width { + rows.push(std::mem::take(&mut row)); + } + row.push_str(cluster); + } + rows.push(row); + rows } /// A horizontal rule: one row of `─`, repeated to the full width. diff --git a/src/markdown/block_tests.rs b/src/markdown/block_tests.rs index 6061d7e..2fc35f7 100644 --- a/src/markdown/block_tests.rs +++ b/src/markdown/block_tests.rs @@ -294,6 +294,51 @@ fn a_code_line_longer_than_the_width_hard_wraps() { ); } +#[test] +fn a_code_line_hard_wraps_at_the_columns_it_prints_in() { + assert_eq!( + rows("```\n日本語\n```", 6), + vec![ + Line::from(Span::styled(" 日本", style::code())), + Line::from(Span::styled(" 語", style::code())), + ] + ); +} + +#[test] +fn a_hard_wrap_keeps_a_letter_and_its_accent_on_one_row() { + assert_eq!( + super::hard_wrap("e\u{301}xy", 2), + vec!["e\u{301}x".to_string(), "y".to_string()] + ); +} + +#[test] +fn a_hard_wrap_at_a_width_of_zero_puts_one_cluster_on_each_row() { + assert_eq!( + super::hard_wrap("ab", 0), + vec!["a".to_string(), "b".to_string()] + ); +} + +#[test] +fn an_html_block_on_a_terminal_with_no_columns_still_renders() { + assert_eq!( + rows("

hi

", 0), + vec![ + Line::raw("<"), + Line::raw("p"), + Line::raw(">"), + Line::raw("h"), + Line::raw("i"), + Line::raw("<"), + Line::raw("/"), + Line::raw("p"), + Line::raw(">"), + ] + ); +} + #[test] fn a_horizontal_rule_fills_the_width() { assert_eq!( diff --git a/src/markdown/table.rs b/src/markdown/table.rs index 5d22b9d..07dda7f 100644 --- a/src/markdown/table.rs +++ b/src/markdown/table.rs @@ -111,25 +111,17 @@ fn natural_widths(header: &[Text<'static>], body: &[Row]) -> Vec { fn content_width(cell: &Text<'static>) -> usize { wrap_spans(cell.clone(), usize::MAX) .iter() - .map(line_width) + .map(Line::width) .max() .unwrap_or(0) } -/// How many characters a row of spans prints. -fn line_width(line: &Line<'static>) -> usize { - line.spans - .iter() - .map(|span| span.content.chars().count()) - .sum() -} - /// The column widths to render at: the natural widths when the table /// fits, otherwise narrower ones. /// /// While the table is too wide, the widest column gives up one -/// character. A run of columns that tie for widest therefore shrinks -/// together, one character at a time each, and a column narrower than +/// column. A run of columns that tie for widest therefore shrinks +/// together, one column at a time each, and a column narrower than /// the widest keeps its width until the widest comes down to meet it. fn fitted(natural: &[usize], width: usize) -> Vec { let mut widths = natural.to_vec(); @@ -205,12 +197,14 @@ fn edge() -> Span<'static> { Span::styled(style::border_set().vertical, style::border()) } -/// One cell's content on one row of it, padded to `width` by its -/// column's alignment, with one more space of padding on each side. A -/// `line` of `None` is a row past the end of a cell that wrapped shorter -/// than its neighbors, and pads to blank. +/// One cell's content on one row of it, padded to `width` columns by its +/// column's alignment, with one more space of padding on each side. The +/// padding counts the columns the content prints in, not its characters, +/// so a cell holding an emoji ends at the same column as the cells above +/// and below it. A `line` of `None` is a row past the end of a cell that +/// wrapped shorter than its neighbors, and pads to blank. fn pad(line: Option<&Line<'static>>, width: usize, alignment: Alignment) -> Vec> { - let slack = width.saturating_sub(line.map_or(0, line_width)); + let slack = width.saturating_sub(line.map_or(0, Line::width)); let (left, right) = match alignment { Alignment::Right => (slack, 0), Alignment::Center => (slack / 2, slack - slack / 2), diff --git a/src/markdown/table_tests.rs b/src/markdown/table_tests.rs index 520031e..6a67b79 100644 --- a/src/markdown/table_tests.rs +++ b/src/markdown/table_tests.rs @@ -192,6 +192,38 @@ fn a_table_too_wide_at_the_minimum_still_renders_every_row() { ); } +#[test] +fn a_cell_holding_a_two_column_emoji_pads_to_its_column_in_columns() { + assert_eq!( + plain("| Name | Icon |\n| --- | --- |\n| Die | 🎲 |", 40), + vec![ + "┌──────┬──────┐", + "│ Name │ Icon │", + "├──────┼──────┤", + "│ Die │ 🎲 │", + "└──────┴──────┘", + ] + ); +} + +#[test] +fn every_row_of_a_table_with_an_emoji_cell_is_the_same_number_of_columns_wide() { + let widths: Vec = rows("| Name | Icon |\n| --- | --- |\n| Die | 🎲 |", 40) + .iter() + .map(Line::width) + .collect(); + + assert_eq!(widths, vec![15, 15, 15, 15, 15]); +} + +#[test] +fn a_column_of_wide_characters_is_as_wide_as_they_print() { + assert_eq!( + plain("| h |\n| --- |\n| 日本 |", 40), + vec!["┌──────┐", "│ h │", "├──────┤", "│ 日本 │", "└──────┘"] + ); +} + #[test] fn a_table_and_a_paragraph_are_separated_by_one_empty_row() { assert_eq!( diff --git a/src/play/editor.rs b/src/play/editor.rs index 693894a..89b210e 100644 --- a/src/play/editor.rs +++ b/src/play/editor.rs @@ -1,4 +1,14 @@ //! A pure text editor: cursor movement and editing across lines, no I/O. +//! +//! Positions inside the buffer are char indices. What the terminal draws +//! is measured in two other units, and both appear here: an edit that +//! removes a symbol removes a whole grapheme cluster, so a letter and its +//! accent go together, and a wrapped row is as long as the columns it +//! prints in, so a row of Japanese breaks at half as many characters as a +//! row of English. + +use unicode_segmentation::UnicodeSegmentation; +use unicode_width::UnicodeWidthStr; /// The prompt's text, which may hold newlines, and where the cursor sits /// in it, a char index from 0 to the character count. @@ -44,27 +54,52 @@ impl Editor { self.cursor += text.chars().count(); } - /// Removes the character before the cursor. Does nothing at the start. + /// Removes the symbol before the cursor, every character of it: a + /// letter takes its accent with it, and an emoji built from several + /// characters goes in one press. Does nothing at the start. pub fn backspace(&mut self) { if self.cursor == 0 { return; } - let start = byte_index(&self.text, self.cursor - 1); + let previous = self.cluster_start_behind(); + let start = byte_index(&self.text, previous); let end = byte_index(&self.text, self.cursor); self.text.replace_range(start..end, ""); - self.cursor -= 1; + self.cursor = previous; } - /// Removes the character under the cursor. Does nothing at the end. + /// Removes the symbol under the cursor, every character of it, the + /// same way [`Editor::backspace`] removes the one behind it. Does + /// nothing at the end. pub fn delete(&mut self) { if self.cursor >= self.len() { return; } let start = byte_index(&self.text, self.cursor); - let end = byte_index(&self.text, self.cursor + 1); + let end = byte_index(&self.text, self.cluster_end_ahead()); self.text.replace_range(start..end, ""); } + /// The char index the symbol in front of the cursor starts at. A + /// cursor sitting inside a cluster, which the character-at-a-time + /// [`Editor::left`] can leave it doing, comes back to that cluster's + /// own start. + fn cluster_start_behind(&self) -> usize { + clusters(&self.text) + .map(|(start, _)| start) + .take_while(|&start| start < self.cursor) + .last() + .expect("a cursor past the start has a cluster starting at 0 behind it") + } + + /// The char index just past the symbol under the cursor. + fn cluster_end_ahead(&self) -> usize { + clusters(&self.text) + .map(|(start, cluster)| start + cluster.chars().count()) + .find(|&end| end > self.cursor) + .expect("a cursor short of the end has a cluster ending past it") + } + /// Moves the cursor one character left. Stops at the start. pub fn left(&mut self) { self.cursor = self.cursor.saturating_sub(1); @@ -144,16 +179,20 @@ impl Editor { /// it. A newline counts as whitespace, so this crosses into the /// previous line when the cursor starts at or near the current /// line's beginning. + /// + /// The walk steps by symbol, not by character, so it lands where a + /// removal can start: a mark riding on a space belongs to the space + /// and moves with it. pub fn word_left(&mut self) { - let chars: Vec = self.text.chars().collect(); - let mut at = self.cursor; - while at > 0 && chars[at - 1].is_whitespace() { + let clusters: Vec<(usize, &str)> = clusters(&self.text).collect(); + let mut at = clusters.partition_point(|&(start, _)| start < self.cursor); + while at > 0 && whitespace(clusters[at - 1].1) { at -= 1; } - while at > 0 && !chars[at - 1].is_whitespace() { + while at > 0 && !whitespace(clusters[at - 1].1) { at -= 1; } - self.cursor = at; + self.cursor = clusters.get(at).map_or(0, |&(start, _)| start); } /// Moves the cursor past the word to its right, emacs-style: past any @@ -257,12 +296,18 @@ impl Editor { row_ranges(&self.text, width).len() } - /// The cursor's row and column among `wrapped_rows(width)`. A cursor - /// sitting on the character a break consumed, a space or a newline, - /// counts as the end of the row before it rather than the start of - /// the row after. + /// The cursor's row among `wrapped_rows(width)` and the screen column + /// it sits at on that row, which is how many columns the row prints + /// in ahead of it. A cursor sitting on the character a break + /// consumed, a space or a newline, counts as the end of the row + /// before it rather than the start of the row after. pub fn cursor_row_column(&self, width: usize) -> (usize, usize) { - locate(&row_ranges(&self.text, width), self.cursor) + let ranges = row_ranges(&self.text, width); + let (row, characters) = locate(&ranges, self.cursor); + let chars: Vec = self.text.chars().collect(); + let (start, _) = ranges[row]; + let ahead: String = chars[start..start + characters].iter().collect(); + (row, ahead.width()) } /// Moves the cursor up one wrapped row at `width` columns, to the @@ -291,6 +336,49 @@ impl Editor { } } +/// `text`'s grapheme clusters, each with the char index it starts at. A +/// cluster is what the terminal draws as one symbol, and what an edit +/// treats as one unit: a letter and the accent behind it are one +/// cluster, and so is an emoji spelled with zero width joiners. +fn clusters(text: &str) -> impl Iterator { + let mut at = 0; + text.graphemes(true).map(move |cluster| { + let start = at; + at += cluster.chars().count(); + (start, cluster) + }) +} + +/// True for a cluster the word walk counts as a gap between words: one +/// whose own character is whitespace, whatever rides on it. +fn whitespace(cluster: &str) -> bool { + cluster.starts_with(char::is_whitespace) +} + +/// One entry per character of `line`: the columns the cluster starting +/// there prints in, and zero for a character that continues the cluster +/// in front of it. +/// +/// Holding a whole cluster's columns on the character that opens it is +/// what keeps a row break out of the middle of one. A continuation +/// character adds nothing to the row it lands on, so the row is never +/// full at one, and the break falls at the next cluster instead. +fn cluster_widths(line: &[char]) -> Vec { + let text: String = line.iter().collect(); + let mut widths = vec![0; line.len()]; + let mut at = 0; + for cluster in text.graphemes(true) { + widths[at] = cluster.width(); + at += cluster.chars().count(); + } + widths +} + +/// The columns a run of characters prints in, from their `cluster_widths`. +fn columns(widths: &[usize]) -> usize { + widths.iter().sum() +} + /// The byte offset of the `chars`-th character, or the end of the string /// once `chars` reaches or passes the last one. fn byte_index(text: &str, chars: usize) -> usize { @@ -300,9 +388,9 @@ fn byte_index(text: &str, chars: usize) -> usize { .unwrap_or(text.len()) } -/// The row and column `cursor` falls on among `ranges`, `row_ranges`'s -/// output: the row whose range reaches at least as far as `cursor`, and -/// how far into it `cursor` sits. +/// The row `cursor` falls on among `ranges`, `row_ranges`'s output, and +/// how many characters into that row it sits: the row whose range reaches +/// at least as far as `cursor`, and the distance from that row's start. fn locate(ranges: &[(usize, usize)], cursor: usize) -> (usize, usize) { let (row, &(start, _)) = ranges .iter() @@ -316,7 +404,8 @@ fn locate(ranges: &[(usize, usize)], cursor: usize) -> (usize, usize) { /// hard newline splits the text into lines first, and each line /// soft-wraps on its own, greedily filling each row and breaking at the /// last whitespace that fits. A line with no whitespace to break on -/// splits mid-word at exactly `width` characters. +/// splits mid-word at the last symbol that leaves the row inside +/// `width` columns. fn row_ranges(text: &str, width: usize) -> Vec<(usize, usize)> { let width = width.max(1); let chars: Vec = text.chars().collect(); @@ -340,6 +429,7 @@ fn wrap_range(line: &[char], offset: usize, width: usize) -> Vec<(usize, usize)> if line.is_empty() { return vec![(offset, offset)]; } + let widths = cluster_widths(line); let mut rows = Vec::new(); let mut row_start = 0; let mut last_space = None; @@ -348,7 +438,7 @@ fn wrap_range(line: &[char], offset: usize, width: usize) -> Vec<(usize, usize)> if line[at].is_whitespace() { last_space = Some(at); } - if at - row_start >= width { + if at > row_start && columns(&widths[row_start..at]) + widths[at] > width { row_start = match last_space { Some(space) if space >= row_start => { rows.push((row_start + offset, space + offset)); diff --git a/src/play/editor_tests.rs b/src/play/editor_tests.rs index 5db826e..3d199ac 100644 --- a/src/play/editor_tests.rs +++ b/src/play/editor_tests.rs @@ -45,6 +45,37 @@ fn backspace_at_the_start_does_nothing() { assert_eq!(editor.cursor(), 0); } +#[test] +fn backspace_removes_a_letter_and_its_accent_together() { + let mut editor = typed("ae\u{301}"); + + editor.backspace(); + + assert_eq!(editor.text(), "a"); + assert_eq!(editor.cursor(), 1); +} + +#[test] +fn backspace_removes_an_emoji_built_from_several_characters_in_one_go() { + let mut editor = typed("a👨\u{200d}👩\u{200d}👧"); + + editor.backspace(); + + assert_eq!(editor.text(), "a"); + assert_eq!(editor.cursor(), 1); +} + +#[test] +fn delete_removes_a_letter_and_its_accent_together() { + let mut editor = typed("e\u{301}a"); + editor.home(); + + editor.delete(); + + assert_eq!(editor.text(), "a"); + assert_eq!(editor.cursor(), 0); +} + #[test] fn delete_removes_the_character_under_the_cursor() { let mut editor = typed("ab"); @@ -211,6 +242,18 @@ fn word_left_skips_whitespace_then_the_word_behind_it() { assert_eq!(editor.cursor(), 0); } +#[test] +fn word_left_over_a_space_carrying_a_combining_mark_stops_past_the_mark() { + // The mark rides on the space, so the two are one symbol of + // whitespace and the word behind the cursor starts at the `c`, the + // fifth character. + let mut editor = typed("ab \u{301}cd"); + + editor.word_left(); + + assert_eq!(editor.cursor(), 4); +} + #[test] fn word_right_from_the_start_goes_to_the_words_end() { let mut editor = typed("ab cd"); diff --git a/src/play/editor_wrap_tests.rs b/src/play/editor_wrap_tests.rs index 433f831..f22e515 100644 --- a/src/play/editor_wrap_tests.rs +++ b/src/play/editor_wrap_tests.rs @@ -127,6 +127,28 @@ fn down_clamps_the_column_to_a_shorter_row() { assert_eq!(editor.cursor(), 6); } +#[test] +fn wrapped_rows_wraps_wide_characters_at_the_columns_they_print_in() { + let editor = typed("日本語"); + + assert_eq!(editor.wrapped_rows(4), vec!["日本", "語"]); +} + +#[test] +fn cursor_row_column_counts_the_columns_wide_characters_print_in() { + let editor = typed("日本語"); + + assert_eq!(editor.cursor_row_column(4), (1, 2)); +} + +#[test] +fn cursor_row_column_after_one_wide_character_is_two_columns_in() { + let mut editor = typed("日本"); + editor.left(); + + assert_eq!(editor.cursor_row_column(40), (0, 2)); +} + #[test] fn wrapped_rows_of_a_long_word_at_width_one_splits_every_character() { let editor = typed("abc"); diff --git a/src/play/prompt.rs b/src/play/prompt.rs index 6fa08fe..f9eb37b 100644 --- a/src/play/prompt.rs +++ b/src/play/prompt.rs @@ -3,6 +3,7 @@ use ratatui::layout::{Position, Rect}; use ratatui::text::{Line, Text}; +use unicode_width::UnicodeWidthStr; use super::editor::Editor; @@ -20,7 +21,7 @@ const CONTINUATION: &str = " "; /// `width` columns wide. The marker and the continuation are the same /// width, so every row of the input wraps at the same column. pub(super) fn text_width(width: u16) -> usize { - (width as usize).saturating_sub(PLAYER_MARKER.len()) + (width as usize).saturating_sub(PLAYER_MARKER.width()) } /// How many rows of `area` the prompt takes: one for each row the input @@ -72,10 +73,13 @@ fn marker(row: usize) -> &'static str { } } -/// Where the cursor goes: over the character at `column` of the prompt's -/// `row`th row, never past the last column. +/// Where the cursor goes: at screen `column` of the prompt's `row`th +/// row, past the marker, never past the last column. `column` counts the +/// columns the row prints in ahead of the cursor, which is what +/// [`Editor::cursor_row_column`] hands back, so a cursor after a wide +/// character lands past all of it. fn cursor(row: usize, column: usize, prompt: Rect) -> Position { - let column = PLAYER_MARKER.len() + column; + let column = PLAYER_MARKER.width() + column; let last = prompt.width.saturating_sub(1); Position::new(prompt.x + (column as u16).min(last), prompt.y + row as u16) } @@ -101,6 +105,32 @@ pub(super) fn on_screen(cursor: Position, area: Rect) -> Position { mod tests { use super::*; + #[test] + fn the_cursor_lands_past_every_column_the_input_prints_in() { + let mut input = Editor::default(); + input.set_text("日本"); + let area = Rect::new(0, 0, 20, 2); + + let (_, cursor) = prompt_widget(&input, area); + + // Two characters, four columns, behind the two-column marker. + assert_eq!(cursor, Position::new(6, 0)); + } + + #[test] + fn an_input_of_wide_characters_wraps_at_the_columns_it_prints_in() { + let mut input = Editor::default(); + input.set_text("日本語"); + let area = Rect::new(0, 0, 6, 2); + + let (text, _) = prompt_widget(&input, area); + + assert_eq!( + text, + Text::from(vec![Line::raw("> 日本"), Line::raw(" 語")]) + ); + } + #[test] fn an_in_bounds_cursor_below_the_top_of_the_screen_passes_through() { let area = Rect::new(0, 20, 80, 4); diff --git a/src/play/transcript.rs b/src/play/transcript.rs index e334abe..00c7216 100644 --- a/src/play/transcript.rs +++ b/src/play/transcript.rs @@ -332,13 +332,17 @@ impl Transcript { /// Queues one row for the transcript, to go out with the rest of the /// pass's rows. /// - /// An empty row is a paragraph break. It does not go out on its own: - /// it only marks that one blank row belongs ahead of the next row - /// that has something on it. A change of kind marks the same thing, - /// which is what separates the blocks of the transcript from each - /// other. + /// A row with no characters on it is a paragraph break. It does not + /// go out on its own: it only marks that one blank row belongs ahead + /// of the next row that has something on it. A change of kind marks + /// the same thing, which is what separates the blocks of the + /// transcript from each other. + /// + /// The test is for characters, not for columns. A row of characters + /// that print in no columns at all, a combining mark that lost its + /// letter, is content the transcript carries rather than a break. fn emit(&mut self, row: Line<'static>, kind: Kind) { - if row.width() == 0 { + if row.iter().all(|span| span.content.is_empty()) { self.pending_blank = true; return; } diff --git a/src/play/transcript_width_tests.rs b/src/play/transcript_width_tests.rs index 5d6cf8c..4822fdd 100644 --- a/src/play/transcript_width_tests.rs +++ b/src/play/transcript_width_tests.rs @@ -135,6 +135,65 @@ fn a_row_too_wide_for_the_terminal_it_lands_on_keeps_every_word() { ); } +/// A paragraph of two-column characters. It prints in twenty-three +/// columns, which [`WIDE`] holds and [`NARROW`] does not, and in fourteen +/// characters, which both of them hold: a wrap that counts characters +/// leaves it too wide for [`NARROW`]. +const WIDE_CHARACTERS: &str = "日本語 の 文字 は 二 列"; + +/// The rows waiting to go out, as the text they carry. +fn queued_rows(transcript: &Transcript) -> Vec { + transcript + .queued + .iter() + .map(|queued| queued.row.to_string()) + .collect() +} + +/// A transcript holding one round of [`WIDE_CHARACTERS`], rendered at +/// [`WIDE`] and then flushed onto a terminal [`NARROW`] columns wide, so +/// its one row has to wrap again to reach the screen whole. +fn wide_round_flushed_narrow() -> Transcript { + let wide = terminal(WIDE); + let narrow = terminal(NARROW); + let mut transcript = Transcript::default(); + transcript.delta(WIDE_CHARACTERS); + transcript.stream(&wide, DEEP_TAIL).unwrap(); + transcript.flush(&narrow).unwrap(); + transcript +} + +#[test] +fn a_row_of_wide_characters_wraps_to_the_columns_the_terminal_has() { + let transcript = wide_round_flushed_narrow(); + + assert_eq!(queued_rows(&transcript), ["日本語 の 文字 は 二", "列"]); +} + +#[test] +fn every_row_of_a_wrapped_wide_row_fits_the_terminal() { + let transcript = wide_round_flushed_narrow(); + + let widest = transcript + .queued + .iter() + .map(|queued| queued.row.width()) + .max(); + assert_eq!(widest, Some(NARROW as usize)); +} + +#[test] +fn a_row_of_zero_width_characters_still_goes_out() { + let wide = terminal(WIDE); + let mut transcript = Transcript::default(); + + transcript + .insert(&wide, "\u{200b}", aside(), Kind::Aside) + .unwrap(); + + assert_eq!(queued_rows(&transcript), ["\u{200b}"]); +} + /// A table of four columns, each already at the narrowest a column /// shrinks to. It renders twenty-five columns wide at any width, so /// [`NARROW`] cannot hold it. diff --git a/src/srd/fetch.rs b/src/srd/fetch.rs index 953e59c..f13f5f5 100644 --- a/src/srd/fetch.rs +++ b/src/srd/fetch.rs @@ -5,6 +5,7 @@ use std::path::{Path, PathBuf}; use std::time::Duration; use flate2::read::GzDecoder; +use serde_yaml_ng::{Mapping, Value}; use sha2::{Digest, Sha256}; use tar::Archive; @@ -97,28 +98,88 @@ impl SrdMeta { } } - /// Renders this record as the `meta.yaml` this pipeline reads and writes. + /// Renders this record as the `meta.yaml` this pipeline reads and + /// writes, starting from `existing` (the file's current contents, when + /// there is one) and rewriting only the fields fetch owns: `version`, + /// and each source entry's `path`, `url`, and pin. Every other key + /// carries over untouched, including a text source's hand-maintained + /// `patches` block: fetch does not know what patches exist and must + /// never invent or drop that record. Source entries are matched by + /// `path`; an entry `existing` doesn't already have gets a fresh one + /// appended, in the same field order fetch has always written. With no + /// `existing` file, this renders a fresh minimal `meta.yaml` from + /// scratch. /// /// The text source is marked `filtered: true`: vendoring keeps only the /// numbered content directories and the top-level licensing files (see /// `should_vendor`), not the full pinned commit. - pub fn render(&self) -> String { - let lines = [ - format!("version: {}", self.version), - "sources:".to_string(), - format!(" - path: {}", self.pdf_path), - format!(" url: {}", self.pdf_url), - format!(" sha256: {}", self.pdf_sha256), - format!(" - path: {}", self.text_path), - format!(" url: {}", self.text_url), - format!(" commit: {}", self.text_commit), - format!(" tarball_sha256: {}", self.text_tarball_sha256), - " filtered: true".to_string(), - ]; - lines.join("\n") + "\n" + pub fn render(&self, existing: Option<&str>) -> String { + let mut top = existing + .and_then(|text| serde_yaml_ng::from_str::(text).ok()) + .and_then(|value| value.as_mapping().cloned()) + .unwrap_or_default(); + + top.insert( + Value::String("version".to_string()), + Value::String(self.version.clone()), + ); + + let mut sources = top + .get("sources") + .and_then(Value::as_sequence) + .cloned() + .unwrap_or_default(); + + let pdf_entry = owned_source_entry(&mut sources, &self.pdf_path); + set_string(pdf_entry, "path", &self.pdf_path); + set_string(pdf_entry, "url", &self.pdf_url); + set_string(pdf_entry, "sha256", &self.pdf_sha256); + + let text_entry = owned_source_entry(&mut sources, &self.text_path); + set_string(text_entry, "path", &self.text_path); + set_string(text_entry, "url", &self.text_url); + set_string(text_entry, "commit", &self.text_commit); + set_string(text_entry, "tarball_sha256", &self.text_tarball_sha256); + text_entry.insert(Value::String("filtered".to_string()), Value::Bool(true)); + + top.insert( + Value::String("sources".to_string()), + Value::Sequence(sources), + ); + + serde_yaml_ng::to_string(&Value::Mapping(top)) + .expect("a mapping of strings and bools always serializes to YAML") } } +/// The entry in `sources` whose `path` field equals `path`, or a freshly +/// appended empty entry when none matches. `SrdMeta::render` uses this to +/// find the mapping it owns a few fields of without disturbing any entry, +/// or any field of that entry, it doesn't recognize. +fn owned_source_entry<'a>(sources: &'a mut Vec, path: &str) -> &'a mut Mapping { + let position = sources.iter().position(|value| { + value + .as_mapping() + .and_then(|mapping| mapping.get("path")) + .and_then(Value::as_str) + == Some(path) + }); + let index = position.unwrap_or_else(|| { + sources.push(Value::Mapping(Mapping::new())); + sources.len() - 1 + }); + sources[index] + .as_mapping_mut() + .expect("owned_source_entry only ever indexes an entry it confirmed is a mapping") +} + +fn set_string(mapping: &mut Mapping, key: &str, value: &str) { + mapping.insert( + Value::String(key.to_string()), + Value::String(value.to_string()), + ); +} + /// The PDF and markdown sources for one SRD version, and the meta record /// that documents their pins. pub struct SrdSources { @@ -148,8 +209,15 @@ pub enum FetchOutcome { pub enum FetchError { Download(String), Io(String), - HashMismatch { expected: String, actual: String }, + HashMismatch { + expected: String, + actual: String, + }, UnsafeEntry(String), + VendoredTreeExists { + destination: String, + meta_path: String, + }, } impl fmt::Display for FetchError { @@ -163,6 +231,15 @@ impl fmt::Display for FetchError { FetchError::UnsafeEntry(name) => { write!(f, "unsafe tarball entry: {name}") } + FetchError::VendoredTreeExists { + destination, + meta_path, + } => write!( + f, + "vendored tree already exists at '{destination}' and its recorded pin does not \ + match. It may hold hand-applied patches recorded in the patches block of \ + '{meta_path}'. Move '{destination}' aside on purpose, then fetch again." + ), } } } @@ -217,13 +294,24 @@ pub fn fetch_all(sources: &SrdSources) -> Result<(FetchOutcome, FetchOutcome), F /// `config.expected_tarball_sha256`, and unpacks it into `config.destination` /// with the tarball's top-level directory stripped. /// -/// If the destination already exists and `config.meta_path` records -/// `config.commit` as its pin, the download is skipped. Otherwise the -/// destination is treated as stale and replaced, and `config.meta_path` is -/// rewritten from `meta`. +/// If the destination already exists, this never removes it. When +/// `config.meta_path` records `config.commit` as its pin, the tree is +/// already up to date and the fetch is skipped. When the pin doesn't +/// match, or can't be found at all, the tree may hold hand-applied patches +/// that only exist on disk, and fetch has no way to tell an intentional +/// edit from a stale download: it refuses and reports an error rather than +/// guessing. Moving the tree aside is the caller's job. Only when the +/// destination doesn't exist yet does fetch download, unpack, and rewrite +/// `config.meta_path`. pub fn fetch_text(config: &TextFetchConfig, meta: &SrdMeta) -> Result { - if config.destination.exists() && meta_pin_matches(&config.meta_path, &config.commit) { - return Ok(FetchOutcome::AlreadyPresent); + if config.destination.exists() { + if meta_pin_matches(&config.meta_path, &config.commit) { + return Ok(FetchOutcome::AlreadyPresent); + } + return Err(FetchError::VendoredTreeExists { + destination: config.destination.display().to_string(), + meta_path: config.meta_path.display().to_string(), + }); } let bytes = download(&config.url)?; @@ -237,29 +325,50 @@ pub fn fetch_text(config: &TextFetchConfig, meta: &SrdMeta) -> Result bool { let Ok(contents) = fs::read_to_string(meta_path) else { return false; }; - let pin_line = format!("commit: {commit}"); - contents.lines().any(|line| line.trim() == pin_line) + let Ok(value) = serde_yaml_ng::from_str::(&contents) else { + return false; + }; + let Some(sources) = value + .as_mapping() + .and_then(|mapping| mapping.get("sources")) + .and_then(Value::as_sequence) + else { + return false; + }; + sources.iter().any(|source| { + source + .as_mapping() + .and_then(|mapping| mapping.get("commit")) + .and_then(Value::as_str) + == Some(commit) + }) } /// Unpacks a gzipped tarball into `destination`, dropping each entry's /// top-level path component and keeping only entries `should_vendor` -/// accepts. Replaces `destination` if it already exists. +/// accepts. +/// +/// Assumes `destination` does not exist yet. `fetch_text` only ever calls +/// this after confirming that, since overwriting an existing vendored tree +/// can destroy hand-applied patches that have no record anywhere but the +/// tree itself. fn unpack_stripped(tarball: &[u8], destination: &Path) -> Result<(), FetchError> { - if destination.exists() { - fs::remove_dir_all(destination)?; - } fs::create_dir_all(destination)?; let mut archive = Archive::new(GzDecoder::new(tarball)); @@ -633,30 +742,75 @@ mod tests { } #[test] - fn srd_meta_renders_the_documented_yaml_shape() { + fn srd_meta_renders_the_documented_yaml_shape_with_no_existing_file() { let meta = fixture_meta(); - let rendered = meta.render(); + let rendered = meta.render(None); assert_eq!( rendered, format!( "version: 5.2.1\n\ sources:\n\ - \x20 - path: sources/SRD_CC_v5.2.1.pdf\n\ - \x20 url: https://example.test/srd.pdf\n\ - \x20 sha256: {pdf_sha256}\n\ - \x20 - path: sources/dnd.srd.5.2.1/\n\ - \x20 url: https://example.test/dnd.srd.5.2.1\n\ - \x20 commit: deadbeef\n\ - \x20 tarball_sha256: {text_tarball_sha256}\n\ - \x20 filtered: true\n", + - path: sources/SRD_CC_v5.2.1.pdf\n\ + \x20 url: https://example.test/srd.pdf\n\ + \x20 sha256: {pdf_sha256}\n\ + - path: sources/dnd.srd.5.2.1/\n\ + \x20 url: https://example.test/dnd.srd.5.2.1\n\ + \x20 commit: deadbeef\n\ + \x20 tarball_sha256: {text_tarball_sha256}\n\ + \x20 filtered: true\n", pdf_sha256 = "a".repeat(64), text_tarball_sha256 = "b".repeat(64), ) ); } + #[test] + fn srd_meta_render_preserves_an_existing_patches_block_and_extra_fields() { + let meta = fixture_meta(); + let existing = "\ + version: 5.2.0\n\ + extra_top_level_field: keep-me\n\ + sources:\n\ + \x20 - path: sources/SRD_CC_v5.2.1.pdf\n\ + \x20 url: https://old.example.test/srd.pdf\n\ + \x20 sha256: oldpdfhash\n\ + \x20 - path: sources/dnd.srd.5.2.1/\n\ + \x20 url: https://example.test/dnd.srd.5.2.1\n\ + \x20 commit: old-commit\n\ + \x20 tarball_sha256: oldtarballhash\n\ + \x20 filtered: true\n\ + \x20 patches:\n\ + \x20 - path: 11_Monsters/Monsters_Each/Winter_Wolf.md\n\ + \x20 reason: file held only the H1 heading, no body\n\ + \x20 repaired_from: 11_Monsters/Monsters_All.md\n"; + + let rendered = meta.render(Some(existing)); + let value: Value = serde_yaml_ng::from_str(&rendered).unwrap(); + let mapping = value.as_mapping().unwrap(); + let sources = mapping.get("sources").unwrap().as_sequence().unwrap(); + let text_entry = sources[1].as_mapping().unwrap(); + let patches = text_entry.get("patches").unwrap().as_sequence().unwrap(); + + assert_eq!( + mapping.get("extra_top_level_field").unwrap().as_str(), + Some("keep-me") + ); + assert_eq!(mapping.get("version").unwrap().as_str(), Some("5.2.1")); + assert_eq!(text_entry.get("commit").unwrap().as_str(), Some("deadbeef")); + assert_eq!(patches.len(), 1); + assert_eq!( + patches[0] + .as_mapping() + .unwrap() + .get("path") + .unwrap() + .as_str(), + Some("11_Monsters/Monsters_Each/Winter_Wolf.md") + ); + } + #[test] fn is_numbered_content_dir_accepts_two_digit_prefixes() { assert!(is_numbered_content_dir("01_Playing_The_Game")); @@ -832,6 +986,64 @@ mod tests { server.join().unwrap(); } + #[test] + fn meta_pin_matches_the_parsed_commit_field() { + let meta_dir = TempDir::new().unwrap(); + let meta_path = meta_dir.path().join("meta.yaml"); + fs::write( + &meta_path, + "sources:\n - path: sources/dnd.srd.5.2.1/\n commit: deadbeef\n", + ) + .unwrap(); + + assert!(meta_pin_matches(&meta_path, "deadbeef")); + } + + #[test] + fn meta_pin_matches_is_false_when_the_file_is_missing() { + let meta_dir = TempDir::new().unwrap(); + + assert!(!meta_pin_matches( + &meta_dir.path().join("meta.yaml"), + "deadbeef" + )); + } + + #[test] + fn meta_pin_matches_is_false_for_invalid_yaml() { + let meta_dir = TempDir::new().unwrap(); + let meta_path = meta_dir.path().join("meta.yaml"); + fs::write(&meta_path, "not: [valid").unwrap(); + + assert!(!meta_pin_matches(&meta_path, "deadbeef")); + } + + #[test] + fn meta_pin_matches_is_false_without_a_sources_list() { + let meta_dir = TempDir::new().unwrap(); + let meta_path = meta_dir.path().join("meta.yaml"); + fs::write(&meta_path, "version: 5.2.1\n").unwrap(); + + assert!(!meta_pin_matches(&meta_path, "deadbeef")); + } + + #[test] + fn meta_pin_matches_ignores_the_pin_text_inside_an_unrelated_value() { + // A block scalar can contain a line that, read on its own, looks + // exactly like a `commit:` pin line. A parser that walked the real + // YAML structure instead of scanning lines of text must not + // mistake that line for the actual `sources[].commit` field. + let meta_dir = TempDir::new().unwrap(); + let meta_path = meta_dir.path().join("meta.yaml"); + fs::write( + &meta_path, + "sources:\n - path: sources/dnd.srd.5.2.1/\n commit: other-pin\n patches:\n - path: a.md\n reason: |\n commit: deadbeef\n", + ) + .unwrap(); + + assert!(!meta_pin_matches(&meta_path, "deadbeef")); + } + #[test] fn fetch_text_skips_download_when_the_recorded_pin_matches() { let destination_dir = TempDir::new().unwrap(); @@ -889,68 +1101,84 @@ mod tests { assert!(!destination.join(".github").exists()); assert!(!destination.join("docs_original").exists()); assert!(!destination.join("SRD-reForged.png").exists()); - assert_eq!(fs::read_to_string(&meta_path).unwrap(), meta.render()); + assert_eq!(fs::read_to_string(&meta_path).unwrap(), meta.render(None)); server.join().unwrap(); } + /// The names and contents of every file under `dir`, recursively, keyed + /// by their path relative to `dir`. Lets a test snapshot a tree before + /// an operation and assert nothing about it changed afterward. + fn snapshot_tree(dir: &Path) -> std::collections::BTreeMap> { + fn walk(dir: &Path, root: &Path, out: &mut std::collections::BTreeMap>) { + for entry in fs::read_dir(dir).unwrap() { + let entry = entry.unwrap(); + let path = entry.path(); + if path.is_dir() { + walk(&path, root, out); + } else { + out.insert( + path.strip_prefix(root).unwrap().to_path_buf(), + fs::read(&path).unwrap(), + ); + } + } + } + let mut out = std::collections::BTreeMap::new(); + walk(dir, dir, &mut out); + out + } + #[test] - fn fetch_text_redownloads_when_the_existing_pin_does_not_match() { - let body = tiny_tar_gz("dnd.srd.5.2.1-deadbeef"); - let expected = sha256_hex(&body); - let (url, server) = serve_once(body); + fn fetch_text_refuses_an_existing_tree_with_a_mismatched_pin() { let destination_dir = TempDir::new().unwrap(); let destination = destination_dir.path().join("dnd.srd.5.2.1"); - fs::create_dir_all(&destination).unwrap(); + fs::create_dir_all(destination.join("07_Spells")).unwrap(); fs::write(destination.join("stale.md"), b"stale").unwrap(); + fs::write(destination.join("07_Spells/stale-spell.md"), b"stale spell").unwrap(); let meta_dir = TempDir::new().unwrap(); let meta_path = meta_dir.path().join("meta.yaml"); - fs::write( - &meta_path, - "sources:\n - path: sources/dnd.srd.5.2.1/\n commit: old-sha\n", - ) - .unwrap(); + let meta_contents = "sources:\n - path: sources/dnd.srd.5.2.1/\n commit: old-sha\n patches:\n - path: a.md\n"; + fs::write(&meta_path, meta_contents).unwrap(); + let tree_before = snapshot_tree(&destination); let config = TextFetchConfig { - url, + url: "http://127.0.0.1:0".to_string(), destination: destination.clone(), - expected_tarball_sha256: expected, + expected_tarball_sha256: "0".repeat(64), commit: "deadbeef".to_string(), meta_path: meta_path.clone(), }; - let meta = fixture_meta(); - let outcome = fetch_text(&config, &meta).unwrap(); + let error = fetch_text(&config, &fixture_meta()).unwrap_err(); - assert_eq!(outcome, FetchOutcome::Downloaded); - assert!(!destination.join("stale.md").exists()); - assert_eq!( - fs::read(destination.join("07_Spells/Fireball.md")).unwrap(), - b"# Fireball\n" + assert!(matches!(error, FetchError::VendoredTreeExists { .. })); + assert!( + error + .to_string() + .contains(&destination.display().to_string()) ); - assert_eq!(fs::read_to_string(&meta_path).unwrap(), meta.render()); - server.join().unwrap(); + assert!(error.to_string().contains("patches")); + assert_eq!(snapshot_tree(&destination), tree_before); + assert_eq!(fs::read_to_string(&meta_path).unwrap(), meta_contents); } #[test] - fn fetch_text_treats_a_missing_meta_file_as_no_recorded_pin() { - let body = tiny_tar_gz("dnd.srd.5.2.1-deadbeef"); - let expected = sha256_hex(&body); - let (url, server) = serve_once(body); + fn fetch_text_refuses_an_existing_tree_with_no_recorded_pin() { let destination_dir = TempDir::new().unwrap(); let destination = destination_dir.path().join("dnd.srd.5.2.1"); fs::create_dir_all(&destination).unwrap(); let meta_dir = TempDir::new().unwrap(); let config = TextFetchConfig { - url, - destination, - expected_tarball_sha256: expected, + url: "http://127.0.0.1:0".to_string(), + destination: destination.clone(), + expected_tarball_sha256: "0".repeat(64), commit: "deadbeef".to_string(), meta_path: meta_dir.path().join("meta.yaml"), }; - let outcome = fetch_text(&config, &fixture_meta()).unwrap(); + let error = fetch_text(&config, &fixture_meta()).unwrap_err(); - assert_eq!(outcome, FetchOutcome::Downloaded); - server.join().unwrap(); + assert!(matches!(error, FetchError::VendoredTreeExists { .. })); + assert!(destination.exists()); } #[test] diff --git a/src/srd/verify/diff.rs b/src/srd/verify/diff.rs index 6d8cef0..c4dd7ee 100644 --- a/src/srd/verify/diff.rs +++ b/src/srd/verify/diff.rs @@ -3,10 +3,22 @@ //! The edit script is a Levenshtein alignment (insert, delete, substitute //! each cost one word) rather than a longest-common-subsequence alignment, //! because a single changed word should read as one substitution, not as a -//! drop immediately followed by an unrelated add. +//! drop immediately followed by an unrelated add. Building that alignment +//! costs one `u32` cell per pair of words still under consideration after +//! trimming, so a pair of streams with little in common (a chapter file +//! paired against the wrong entry, say) can demand an allocation far +//! larger than the words involved would suggest; `first_divergence` falls +//! back to a cheaper, coarser report rather than let that grow unbounded. const CONTEXT_WORDS: usize = 3; +/// The most alignment-matrix cells `first_divergence` will build before +/// falling back to `bounded_divergence`. At 4 bytes per `u32` cell, this +/// caps the matrix at 32 MB. A 68,000-word chapter file mismatched against +/// a 10,000-word entry, the shape of input that has triggered this before, +/// would otherwise ask for roughly 680,000,000 cells: 2.7 GB. +const MAX_ALIGNMENT_CELLS: usize = 8_000_000; + enum Op { Equal(String), Delete(String), @@ -40,7 +52,13 @@ fn template(op: &Op) -> Option { /// matrix. On a mismatch, the common prefix and common suffix are trimmed /// off both streams first, so the matrix only covers the words that /// actually differ; the trimmed words still supply context when the -/// divergence sits too close to either end of what remains. +/// divergence sits too close to either end of what remains. When even the +/// trimmed streams are too large for the matrix to stay under +/// `MAX_ALIGNMENT_CELLS`, this reports the first position where the +/// trimmed streams disagree instead of computing the full alignment: a +/// coarser report (an insert or delete part way through can read as a run +/// of substitutions), but one that costs no more memory than the streams +/// themselves. pub fn first_divergence(source: &[String], corpus: &[String]) -> Option { if source == corpus { return None; @@ -53,6 +71,15 @@ pub fn first_divergence(source: &[String], corpus: &[String]) -> Option let source_middle = &source[prefix_len..source.len() - suffix_len]; let corpus_middle = &corpus[prefix_len..corpus.len() - suffix_len]; + if alignment_cells(source_middle.len(), corpus_middle.len()) > MAX_ALIGNMENT_CELLS { + return Some(bounded_divergence( + source_middle, + corpus_middle, + prefix, + suffix, + )); + } + let ops = edit_script(source_middle, corpus_middle); let (index, message) = ops .iter() @@ -64,6 +91,89 @@ pub fn first_divergence(source: &[String], corpus: &[String]) -> Option Some(format!("{message} (context: {context})")) } +/// How many cells an alignment matrix over streams of these lengths would +/// need: one row per source word plus the empty prefix, one column per +/// corpus word plus the empty prefix. +fn alignment_cells(source_len: usize, corpus_len: usize) -> usize { + (source_len + 1).saturating_mul(corpus_len + 1) +} + +/// Reports the first position where two already-trimmed word streams +/// disagree, comparing them word by word instead of computing an +/// alignment. Because nothing here realigns the streams after a mismatch, +/// an insert or delete part way through reads as a run of substitutions +/// rather than as the single insert or delete `edit_script` would find; +/// `first_divergence` only reaches this when the streams are too large for +/// that better report to be worth its memory. +fn bounded_divergence( + source_middle: &[String], + corpus_middle: &[String], + prefix: &[String], + suffix: &[String], +) -> String { + let len = source_middle.len().max(corpus_middle.len()); + let index = (0..len) + .find(|&i| source_middle.get(i) != corpus_middle.get(i)) + .expect("first_divergence only calls this with streams already confirmed unequal"); + + let source_word = source_middle.get(index); + let corpus_word = corpus_middle.get(index); + let message = if let (Some(source_word), Some(corpus_word)) = (source_word, corpus_word) { + format!("source has \"{source_word}\" but the corpus body has \"{corpus_word}\"") + } else if let Some(source_word) = source_word { + format!("source has \"{source_word}\" but the corpus body drops it") + } else { + let corpus_word = corpus_word.expect("index is within source_middle or corpus_middle"); + format!("corpus body adds \"{corpus_word}\", which the source does not have") + }; + + let context_stream = if index < source_middle.len() { + source_middle + } else { + corpus_middle + }; + let before = bounded_context_before(context_stream, index, prefix); + let after = bounded_context_after(context_stream, index, suffix); + let context = format!("{before} ... {after}").trim().to_string(); + format!("{message} (context: {context})") +} + +/// The last few words of `stream` before `index`, in reading order, filled +/// in from `prefix` when `stream` alone doesn't hold enough. Mirrors +/// `context_before`, but reads a plain word slice instead of `Op`s. +fn bounded_context_before(stream: &[String], index: usize, prefix: &[String]) -> String { + let mut words: Vec<&str> = stream[..index] + .iter() + .rev() + .take(CONTEXT_WORDS) + .map(String::as_str) + .collect(); + if words.len() < CONTEXT_WORDS { + let remaining = CONTEXT_WORDS - words.len(); + words.extend(prefix.iter().rev().take(remaining).map(String::as_str)); + } + words.reverse(); + words.join(" ") +} + +/// The next few words of `stream` after `index`, in reading order, filled +/// in from `suffix` when `stream` alone doesn't hold enough. Mirrors +/// `context_after`, but reads a plain word slice instead of `Op`s. +fn bounded_context_after(stream: &[String], index: usize, suffix: &[String]) -> String { + let mut words: Vec<&str> = stream + .get(index + 1..) + .unwrap_or(&[]) + .iter() + .take(CONTEXT_WORDS) + .map(String::as_str) + .collect(); + if words.len() < CONTEXT_WORDS { + let remaining = CONTEXT_WORDS - words.len(); + words.extend(suffix.iter().take(remaining).map(String::as_str)); + } + words.join(" ") +} + /// How many leading words `source` and `corpus` share. fn common_prefix_len(source: &[String], corpus: &[String]) -> usize { source @@ -282,4 +392,71 @@ mod tests { first_divergence(&words("roar into fire"), &words("roars into fire")).unwrap(); assert!(message.contains("context: ... into fire")); } + + #[test] + fn alignment_cells_counts_rows_times_columns_including_the_empty_prefix() { + assert_eq!(alignment_cells(2, 3), 3 * 4); + } + + #[test] + fn a_small_entirely_mismatched_pair_stays_under_the_ceiling() { + // A stream this size, entirely mismatched, keeps the trimmed + // middle's cell count comfortably under the ceiling, so this still + // takes the full alignment path. Pairs this size (or the existing + // under-500-word tests above) are the "existing nicer report" + // behavior the ceiling must leave unchanged. + let source: Vec = (0..50).map(|i| format!("s{i}")).collect(); + let corpus: Vec = (0..50).map(|i| format!("c{i}")).collect(); + + assert!(alignment_cells(source.len(), corpus.len()) <= MAX_ALIGNMENT_CELLS); + let message = first_divergence(&source, &corpus).unwrap(); + assert!(message.contains("\"s0\"")); + assert!(message.contains("\"c0\"")); + } + + #[test] + fn a_pair_of_streams_over_the_alignment_ceiling_uses_the_bounded_report() { + // Every word differs and neither end matches, so prefix/suffix + // trimming can't shrink the middle: the whole streams count + // toward the cell ceiling, the same shape of input a mismatched + // `source:` field produces on a real chapter file. + let side = (MAX_ALIGNMENT_CELLS as f64).sqrt() as usize + 10; + let source: Vec = vec!["a".to_string(); side]; + let corpus: Vec = vec!["b".to_string(); side]; + assert!(alignment_cells(source.len(), corpus.len()) > MAX_ALIGNMENT_CELLS); + + let message = first_divergence(&source, &corpus).unwrap(); + + assert!(message.contains("source has \"a\" but the corpus body has \"b\"")); + } + + #[test] + fn bounded_divergence_reports_a_substitution() { + let message = bounded_divergence(&words("a b c"), &words("a x c"), &[], &[]); + assert!(message.contains("source has \"b\" but the corpus body has \"x\"")); + assert!(message.contains("context: a ... c")); + } + + #[test] + fn bounded_divergence_reports_a_source_only_tail() { + let message = bounded_divergence(&words("a b"), &words("a"), &[], &[]); + assert!(message.contains("source has \"b\" but the corpus body drops it")); + } + + #[test] + fn bounded_divergence_reports_a_corpus_only_tail() { + let message = bounded_divergence(&words("a"), &words("a b"), &[], &[]); + assert!(message.contains("corpus body adds \"b\", which the source does not have")); + } + + #[test] + fn bounded_divergence_fills_in_context_from_the_trimmed_prefix_and_suffix() { + let message = bounded_divergence( + &words("roar"), + &words("roars"), + &words("well into the"), + &words("fire indeed"), + ); + assert!(message.contains("context: well into the ... fire indeed")); + } } diff --git a/src/wrap.rs b/src/wrap.rs index 5eb3451..8fdf8e9 100644 --- a/src/wrap.rs +++ b/src/wrap.rs @@ -2,15 +2,23 @@ use ratatui::style::Style; use ratatui::text::{Line, Span, Text}; +use unicode_segmentation::UnicodeSegmentation; +use unicode_width::UnicodeWidthStr; -/// Wraps `text` to `width` characters, one output line per row of the +/// Wraps `text` to `width` columns, one output line per row of the /// terminal. /// +/// A row is measured in the columns the terminal draws it in, not in +/// characters: an emoji takes two columns, a letter and the accent behind +/// it take one between them. This is the same measure +/// [`ratatui::buffer::Buffer::set_line`] cuts a row at, so a row that +/// fits here reaches the screen whole. +/// /// Runs of whitespace collapse to a single space. A newline always starts a /// new row, so blank lines between paragraphs survive. A word longer than -/// `width` breaks across rows. A `width` of zero wraps at one character, -/// which keeps the wrap from looping forever on a terminal that reports no -/// columns. +/// `width` breaks across rows, at a grapheme cluster boundary. A `width` of +/// zero wraps at one column, which keeps the wrap from looping forever on a +/// terminal that reports no columns. pub fn wrap(text: &str, width: usize) -> Vec { let width = width.max(1); text.split('\n') @@ -26,7 +34,7 @@ fn wrap_line(line: &str, width: usize) -> Vec { for piece in break_word(word, width) { if row.is_empty() { row = piece; - } else if row.chars().count() + 1 + piece.chars().count() <= width { + } else if row.width() + 1 + piece.width() <= width { row.push(' '); row.push_str(&piece); } else { @@ -38,16 +46,19 @@ fn wrap_line(line: &str, width: usize) -> Vec { rows } -/// Cuts `word` into pieces of at most `width` characters. A word that fits -/// comes back whole. +/// Cuts `word` into pieces of at most `width` columns. A word that fits +/// comes back whole. A cut lands between two grapheme clusters, never +/// inside one, so an accent stays on its letter and the characters that +/// build one emoji stay together. A cluster wider than `width` is a piece +/// of its own and overflows the row rather than disappearing. fn break_word(word: &str, width: usize) -> Vec { let mut pieces = Vec::new(); let mut piece = String::new(); - for character in word.chars() { - if piece.chars().count() == width { + for cluster in word.graphemes(true) { + if !piece.is_empty() && piece.width() + cluster.width() > width { pieces.push(std::mem::take(&mut piece)); } - piece.push(character); + piece.push_str(cluster); } pieces.push(piece); pieces @@ -99,7 +110,7 @@ fn wrap_styled_line(line: &Line<'static>, width: usize) -> Vec> { for piece in break_styled_word(word.characters, width) { if row.is_empty() { row = piece; - } else if row.len() + 1 + piece.len() <= width { + } else if columns(&row) + 1 + columns(&piece) <= width { row.push((' ', word.lead)); row.extend(piece); } else { @@ -156,21 +167,46 @@ fn styled_words(line: &Line<'static>) -> Vec { words } -/// Cuts `word` into pieces of at most `width` characters, the same way +/// Cuts `word` into pieces of at most `width` columns, the same way /// [`break_word`] cuts plain text. A word that fits comes back whole. fn break_styled_word(word: Vec, width: usize) -> Vec> { let mut pieces = Vec::new(); - let mut piece = Vec::new(); - for character in word { - if piece.len() == width { + let mut piece: Vec = Vec::new(); + for cluster in styled_clusters(&word) { + if !piece.is_empty() && columns(&piece) + columns(cluster) > width { pieces.push(std::mem::take(&mut piece)); } - piece.push(character); + piece.extend_from_slice(cluster); } pieces.push(piece); pieces } +/// The columns a row of styled characters prints in. +fn columns(row: &[StyledChar]) -> usize { + text_of(row).width() +} + +/// `word` cut at its grapheme cluster boundaries, one slice per cluster. +/// A cut between two pieces of a word falls on one of these boundaries, +/// so the characters that draw as a single symbol stay on one row. +fn styled_clusters(word: &[StyledChar]) -> Vec<&[StyledChar]> { + let text = text_of(word); + let mut clusters = Vec::new(); + let mut at = 0; + for cluster in text.graphemes(true) { + let characters = cluster.chars().count(); + clusters.push(&word[at..at + characters]); + at += characters; + } + clusters +} + +/// The characters of a row of styled characters, without their styles. +fn text_of(row: &[StyledChar]) -> String { + row.iter().map(|(character, _)| *character).collect() +} + /// Turns a row of styled characters into a line, one span per run of /// characters that share a style. fn styled_line(row: Vec) -> Line<'static> { @@ -259,6 +295,22 @@ mod tests { assert_eq!(wrap("ab", 0), vec!["a", "b"]); } + #[test] + fn a_two_column_emoji_takes_two_of_the_width() { + // Seven characters, eight columns: the row does not fit. + assert_eq!(wrap("🎲 rolls", 7), vec!["🎲", "rolls"]); + } + + #[test] + fn a_word_of_wide_characters_breaks_by_columns() { + assert_eq!(wrap("日本語", 4), vec!["日本", "語"]); + } + + #[test] + fn a_break_inside_a_word_keeps_a_letter_and_its_accent_together() { + assert_eq!(wrap("e\u{301}xy", 2), vec!["e\u{301}x", "y"]); + } + #[test] fn a_fitting_line_of_spans_is_left_alone() { let dim = Style::new().add_modifier(Modifier::DIM); @@ -357,6 +409,24 @@ mod tests { ); } + #[test] + fn a_span_wrap_counts_a_two_column_emoji_as_two() { + let text = Text::from(Line::from(Span::raw("🎲 rolls"))); + + assert_eq!( + wrap_spans(text, 7), + vec![Line::raw("🎲"), Line::raw("rolls")] + ); + } + + #[test] + fn a_span_wrap_never_splits_an_emoji_built_from_several_characters() { + let family = "👨\u{200d}👩\u{200d}👧"; + let text = Text::from(Line::from(Span::raw(format!("{family}x")))); + + assert_eq!(wrap_spans(text, 2), vec![Line::raw(family), Line::raw("x")]); + } + #[test] fn a_hard_split_keeps_its_span_style() { let dim = Style::new().add_modifier(Modifier::DIM);