From 8e5c17b5543e0e79b1b74d9723f808861f70c162 Mon Sep 17 00:00:00 2001 From: sigoden Date: Thu, 11 Jul 2024 12:25:30 +0800 Subject: refactor: improve estimate_token_length (#703) --- src/client/common.rs | 32 +++++++++++++++++--------------- 1 file changed, 17 insertions(+), 15 deletions(-) (limited to 'src/client/common.rs') diff --git a/src/client/common.rs b/src/client/common.rs index f25560d..725f715 100644 --- a/src/client/common.rs +++ b/src/client/common.rs @@ -368,7 +368,7 @@ pub trait Client: Sync + Send { ret = async { if self.global_config().read().dry_run { let content = input.echo_messages(); - let tokens = tokenize(&content); + let tokens = split_content(&content); for token in tokens { tokio::time::sleep(Duration::from_millis(10)).await; handler.text(token)?; @@ -648,22 +648,22 @@ pub fn catch_error(data: &Value, status: u16) -> Result<()> { debug!("Invalid response, status: {status}, data: {data}"); if let Some(error) = data["error"].as_object() { if let (Some(typ), Some(message)) = ( - get_str_field_from_json_map(error, "type"), - get_str_field_from_json_map(error, "message"), + json_str_from_map(error, "type"), + json_str_from_map(error, "message"), ) { bail!("{message} (type: {typ})"); } } else if let Some(error) = data["errors"][0].as_object() { if let (Some(code), Some(message)) = ( - get_u64_field_from_json_map(error, "code"), - get_str_field_from_json_map(error, "message"), + error.get("code").and_then(|v| v.as_u64()), + json_str_from_map(error, "message"), ) { bail!("{message} (status: {code})") } } else if let Some(error) = data[0]["error"].as_object() { if let (Some(status), Some(message)) = ( - get_str_field_from_json_map(error, "status"), - get_str_field_from_json_map(error, "message"), + json_str_from_map(error, "status"), + json_str_from_map(error, "message"), ) { bail!("{message} (status: {status})") } @@ -678,20 +678,13 @@ pub fn catch_error(data: &Value, status: u16) -> Result<()> { bail!("Invalid response data: {data} (status: {status})"); } -pub fn get_str_field_from_json_map<'a>( +pub fn json_str_from_map<'a>( map: &'a serde_json::Map, field_name: &str, ) -> Option<&'a str> { map.get(field_name).and_then(|v| v.as_str()) } -pub fn get_u64_field_from_json_map( - map: &serde_json::Map, - field_name: &str, -) -> Option { - map.get(field_name).and_then(|v| v.as_u64()) -} - pub fn maybe_catch_error(data: &Value) -> Result<()> { if let (Some(code), Some(message)) = (data["code"].as_str(), data["message"].as_str()) { debug!("Invalid response: {}", data); @@ -768,3 +761,12 @@ fn to_json(kind: &PromptKind, value: &str) -> Value { }, } } + +fn split_content(text: &str) -> Vec<&str> { + if text.is_ascii() { + text.split_inclusive(|c: char| c.is_ascii_whitespace()) + .collect() + } else { + unicode_segmentation::UnicodeSegmentation::graphemes(text, true).collect() + } +} -- cgit v1.2.3