prompt.rs 5.3 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189
  1. use actix_web::{HttpResponse, web::Bytes};
  2. use reqwest::Client;
  3. use serde::Serialize;
  4. use serde_json::{Value, json, from_str};
  5. use futures_util::{StreamExt, stream};
  6. use tokio::sync::oneshot;
  7. use chrono::Utc;
  8. use std::sync::{Arc, Mutex};
  9. use crate::{
  10. http_error::HttpError,
  11. models::{
  12. game::Game,
  13. character::Character,
  14. session::Session,
  15. turn::Turn
  16. },
  17. DEEPINFRA_TOKEN
  18. };
  19. #[derive(Serialize)]
  20. struct PromptMessage {
  21. pub role: PromptRole,
  22. pub content: String
  23. }
  24. #[derive(Serialize)]
  25. #[serde(rename_all = "lowercase")]
  26. enum PromptRole {
  27. User,
  28. System,
  29. Assistant
  30. }
  31. pub struct PromptResponse {
  32. pub message: oneshot::Receiver<(String, i32, i32)>,
  33. pub sse: HttpResponse
  34. }
  35. pub async fn prompt(
  36. game: Game,
  37. character: Character,
  38. sessions: Vec<Session>,
  39. turns: Vec<Turn>,
  40. new_input: Option<String>,
  41. new_game: bool
  42. ) -> Result<PromptResponse, HttpError> {
  43. let system_message = create_system_message(game, character, sessions, new_game);
  44. let mut prompt = create_turn_prompts(turns);
  45. prompt.insert(0, system_message);
  46. if let Some(p) = new_input {
  47. prompt.push(PromptMessage {
  48. role: PromptRole::User,
  49. content: p
  50. });
  51. }
  52. send_message(prompt).await
  53. }
  54. fn create_system_message(
  55. game: Game,
  56. character: Character,
  57. sessions: Vec<Session>,
  58. new_game: bool
  59. ) -> PromptMessage {
  60. let mut text = "".to_string();
  61. text += "You are acting as a game master for an RPG and running the game.\n";
  62. text += &format!("Here is the world context: {}\n", game.game_context);
  63. text += &format!(
  64. "Character name is {}, here is a description of the character: {}\n",
  65. character.name,
  66. character.description
  67. );
  68. if sessions.len() == 0 {
  69. text += "This is the first session in this game.\n";
  70. } else {
  71. for(i, session) in sessions.iter().enumerate() {
  72. match &session.summary {
  73. Some(s) => {text += &format!("Summary of session {}: {}\n", i, s)},
  74. None => ()
  75. }
  76. }
  77. }
  78. if new_game {
  79. text += "Start off a new session for the player";
  80. }
  81. PromptMessage {
  82. role: PromptRole::System,
  83. content: text
  84. }
  85. }
  86. fn create_turn_prompts(turns: Vec<Turn>) -> Vec<PromptMessage> {
  87. let mut prompts = Vec::new();
  88. for turn in turns {
  89. prompts.push(PromptMessage {
  90. role: PromptRole::User,
  91. content: turn.user_text
  92. });
  93. prompts.push(PromptMessage {
  94. role: PromptRole::Assistant,
  95. content: turn.llm_text
  96. });
  97. }
  98. prompts
  99. }
  100. async fn send_message(prompt: Vec<PromptMessage>) -> Result<PromptResponse, HttpError> {
  101. let client = Client::new();
  102. let key = DEEPINFRA_TOKEN.get().unwrap();
  103. let body = json!({
  104. "model": "Qwen/Qwen3-32B",
  105. "messages": prompt,
  106. "reasoning_effort": "none",
  107. "temperature": 0.7,
  108. "stream": true,
  109. "stream_options": {"include_usage": true},
  110. "max_tokens": 4096
  111. });
  112. let response = client
  113. .post("https://api.deepinfra.com/v1/chat/completions")
  114. .header("Authorization", format!("Bearer {}", key))
  115. .header("Content-Type", "application/json")
  116. .json(&body)
  117. .send()
  118. .await
  119. .map_err(|e| HttpError::InternalError(e.to_string()))?;
  120. let (tx, rx) = oneshot::channel::<(String, i32, i32)>();
  121. let full_text = Arc::new(Mutex::new(String::new()));
  122. let acc = full_text.clone();
  123. let stream = response.bytes_stream().map(move |chunk| -> Result<Bytes, HttpError> {
  124. match chunk {
  125. Ok(bytes) => {
  126. acc.lock().unwrap().push_str(&String::from_utf8_lossy(&bytes));
  127. Ok(bytes)
  128. },
  129. Err(e) => Err(HttpError::InternalError(e.to_string()))
  130. }
  131. }).chain(stream::once(async move {
  132. let raw = full_text.lock().unwrap().clone();
  133. let content = raw.lines()
  134. .filter_map(|l| l.strip_prefix("data: "))
  135. .filter(|d| *d != "[DONE]")
  136. .filter_map(|d| from_str::<Value>(d).ok())
  137. .filter_map(|v| v["choices"][0]["delta"]["content"].as_str().map(str::to_string))
  138. .collect::<String>();
  139. let (input_tokens, output_tokens) = raw.lines()
  140. .filter_map(|l| l.strip_prefix("data: "))
  141. .filter_map(|d| from_str::<Value>(d).ok())
  142. .find_map(|v| {
  143. let usage = v.get("usage")?;
  144. if usage.is_null() { return None; }
  145. Some(usage.clone())
  146. })
  147. .map(|u| (
  148. u["prompt_tokens"].as_i64().unwrap_or(0) as i32,
  149. u["completion_tokens"].as_i64().unwrap_or(0) as i32
  150. ))
  151. .unwrap_or((0, 0));
  152. let meta = json!({
  153. "tokens": input_tokens + (output_tokens * 2),
  154. "created_at": Utc::now(),
  155. "something": "else"
  156. });
  157. let sse_chunk = format!("event: meta\ndata: {}\n\n", meta.to_string());
  158. let _ = tx.send((content, input_tokens, output_tokens));
  159. Ok::<Bytes, HttpError>(Bytes::from(sse_chunk))
  160. }));
  161. Ok(PromptResponse {
  162. message: rx,
  163. sse: HttpResponse::Ok()
  164. .content_type("text/event-stream")
  165. .streaming(stream)
  166. })
  167. }