appchat.py 19 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436437438439440441442443444445446447448449450451452453454455456457458459460461462463464465466467468469470471472473474475476477478479480481482483484485486487488489490491492493494495496497498499500501502503504505506507508509510511512513514515516517518519520521522523524
  1. # -*- coding: utf-8 -*-
  2. import os
  3. import re
  4. import glob
  5. import json
  6. import shutil
  7. import base64
  8. import uuid
  9. import datetime
  10. import subprocess
  11. import html
  12. import streamlit as st
  13. import streamlit.components.v1 as components
  14. from openai import OpenAI
  15. # ========== 启动文件服务(非阻塞)==========
  16. try:
  17. if "fastpdf_spawned" not in st.session_state:
  18. subprocess.Popen(["python", "fastpdf.py"])
  19. st.session_state["fastpdf_spawned"] = True
  20. except Exception:
  21. pass
  22. # ========== 示例:模型和相应的 API 密钥 ==========
  23. # 请在生产环境中改为安全地注入 API Key(如通过环境变量)
  24. default_key = "sk-re2NlaKIQn11ZNWzAbB6339cEbF94c6aAfC8B7Ab82879bEa" # 默认令牌,请勿用于生产
  25. model_keys = {
  26. "grok-3": default_key,
  27. "grok-4.1": default_key,
  28. "grok-4.2": default_key,
  29. "gpt-5.4-mini": default_key,
  30. "gpt-4o-mini": default_key,
  31. # "gpt-4.1-mini-2025-04-14": default_key,
  32. "o1-mini": default_key,
  33. "o4-mini": default_key,
  34. "deepseek-v3.1": default_key,
  35. "deepseek-r1": default_key,
  36. "grok-4": default_key,
  37. "gpt-5-all": default_key,
  38. "gpt-4o-all": default_key,
  39. # "gpt-5-mini-2025-08-07": default_key,
  40. "o3-mini-all": default_key,
  41. # "claude-sonnet-4-20250514": default_key,
  42. }
  43. st.title("ChatGPT-like Clone")
  44. selected_model = st.sidebar.selectbox("选择模型", list(model_keys.keys()))
  45. api_key = model_keys[selected_model]
  46. api_url = "https://yunwu.ai/v1"
  47. client = OpenAI(api_key=api_key, base_url=api_url)
  48. data_dir = "data"
  49. backup_dir = "data_bak"
  50. blog_dir = "blog"
  51. upload_dir = "uploads"
  52. for directory in [data_dir, backup_dir, blog_dir, upload_dir]:
  53. if not os.path.exists(directory):
  54. os.makedirs(directory)
  55. session_id_file = os.path.join(data_dir, "session_id.txt")
  56. def load_session_id():
  57. if os.path.exists(session_id_file):
  58. try:
  59. with open(session_id_file, "r", encoding="utf-8") as f:
  60. c = f.read().strip()
  61. return int(c) if c.isdigit() else 0
  62. except Exception:
  63. return 0
  64. return 0
  65. def save_session_id(session_id: int):
  66. with open(session_id_file, "w", encoding="utf-8") as f:
  67. f.write(str(session_id))
  68. if "openai_model" not in st.session_state:
  69. st.session_state["openai_model"] = selected_model
  70. def load_history(file_path: str):
  71. match = re.search(r'chat_history_(\d+)\.json', os.path.basename(file_path))
  72. if match:
  73. st.session_state.session_id = int(match.group(1))
  74. try:
  75. with open(file_path, "r", encoding="utf-8") as f:
  76. st.session_state.messages = json.load(f)
  77. except Exception:
  78. st.session_state.messages = []
  79. def load_latest_history():
  80. history_files = sorted(glob.glob(os.path.join(data_dir, "*.json")), key=os.path.getmtime)
  81. if history_files:
  82. load_history(history_files[-1])
  83. if "messages" not in st.session_state:
  84. st.session_state.messages = []
  85. load_latest_history()
  86. if "session_id" not in st.session_state:
  87. st.session_state.session_id = load_session_id()
  88. # 用于控制“单块展开”的状态集合(索引集合)
  89. if "expanded_messages" not in st.session_state:
  90. st.session_state.expanded_messages = set()
  91. def save_history():
  92. file_path = os.path.join(data_dir, f"chat_history_{st.session_state.session_id}.json")
  93. try:
  94. json.dumps(st.session_state.messages)
  95. with open(file_path, "w", encoding="utf-8") as f:
  96. json.dump(st.session_state.messages, f, ensure_ascii=False)
  97. except (TypeError, ValueError):
  98. st.error("数据格式不正确,无法保存记录")
  99. def move_history(file_path: str):
  100. if os.path.exists(file_path):
  101. destination_file = os.path.join(backup_dir, os.path.basename(file_path))
  102. shutil.move(file_path, destination_file)
  103. def delete_history(file_path: str):
  104. if os.path.exists(file_path):
  105. os.remove(file_path)
  106. def text_from_content(content):
  107. if isinstance(content, str):
  108. return content
  109. elif isinstance(content, list):
  110. return " ".join([part.get("text", "") for part in content if part.get("type") == "text"])
  111. return str(content)
  112. def get_chat_title(messages):
  113. if messages:
  114. title = text_from_content(messages[0].get("content", "")) or "空的聊天"
  115. return title[:12]
  116. return "空的聊天"[:12]
  117. def show_history_files(page: int = 0, page_size: int = 10):
  118. with st.sidebar.expander("历史聊天记录"):
  119. history_files = [(f, os.path.getmtime(f)) for f in glob.glob(os.path.join(data_dir, "*.json"))]
  120. history_files.sort(key=lambda x: x[1], reverse=True)
  121. total_files = len(history_files)
  122. total_pages = (total_files // page_size) + (1 if total_files % page_size > 0 else 0)
  123. # 页码范围校正
  124. if page >= total_pages and total_pages > 0:
  125. st.session_state.current_page = max(0, total_pages - 1)
  126. st.rerun()
  127. elif page < 0:
  128. st.session_state.current_page = 0
  129. st.rerun()
  130. start = page * page_size
  131. end = start + page_size
  132. display_files = history_files[start:end]
  133. if not display_files:
  134. st.write("无记录。")
  135. else:
  136. for file_path, _mtime in display_files:
  137. try:
  138. with open(file_path, "r", encoding="utf-8") as f:
  139. chat_history = json.load(f)
  140. except Exception:
  141. continue
  142. file_name = get_chat_title(chat_history)
  143. col1, col2, col3 = st.sidebar.columns([3, 1, 1])
  144. with col1:
  145. if st.button(file_name, key=f"load_{os.path.basename(file_path)}"):
  146. load_history(file_path)
  147. # 切换历史时,清空单块展开状态
  148. st.session_state.expanded_messages = set()
  149. st.rerun()
  150. with col2:
  151. if st.button("📦", key=f"move_{os.path.basename(file_path)}", help="移动到备份文件夹"):
  152. move_history(file_path)
  153. st.success(f"{file_name} 已移动到备份。")
  154. st.rerun()
  155. with col3:
  156. if st.button("❌", key=f"delete_{os.path.basename(file_path)}", help="删除"):
  157. delete_history(file_path)
  158. st.success(f"{file_name} 已删除。")
  159. st.rerun()
  160. col_prev, col_next = st.sidebar.columns([1, 1])
  161. with col_prev:
  162. if st.button("上一页", disabled=page <= 0, key="history_prev_page"):
  163. st.session_state.current_page -= 1
  164. st.rerun()
  165. with col_next:
  166. if st.button("下一页", disabled=page >= total_pages - 1, key="history_next_page"):
  167. st.session_state.current_page += 1
  168. st.rerun()
  169. if "current_page" not in st.session_state:
  170. st.session_state.current_page = 0
  171. def export_message(content):
  172. # 文本提取与转义
  173. if isinstance(content, str):
  174. processed_content = content.replace('\n', '\n').replace('"', '"')
  175. elif isinstance(content, list):
  176. processed_content = " ".join([part.get("text", "") for part in content if part.get("type") == "text"])
  177. processed_content = processed_content.replace('\n', '\n').replace('"', '"')
  178. else:
  179. processed_content = str(content)
  180. timestamp = datetime.datetime.now().strftime("%m%d%H%M")
  181. first_10 = (
  182. processed_content[:10]
  183. .replace(' ', '')
  184. .replace('/', '')
  185. .replace('\\', '')
  186. .replace(':', '')
  187. .replace('```', '')
  188. )
  189. filename = f"{timestamp}_{first_10}.txt"
  190. file_path = os.path.join(blog_dir, filename)
  191. try:
  192. with open(file_path, "w", encoding="utf-8") as f:
  193. f.write(processed_content)
  194. st.success(f"已导出到 {file_path}")
  195. except Exception as e:
  196. st.error(f"导出失败:{e}")
  197. # 侧边栏:输出模式
  198. output_mode = st.sidebar.selectbox("选择输出模式", ["流式输出 (Stream)", "非流式输出 (Non-stream)"])
  199. # ========== 高亮样式与收缩样式 ==========
  200. st.markdown("""
  201. <style>
  202. mark.hl {
  203. background: #fff3a3;
  204. color: inherit;
  205. padding: 0 2px;
  206. border-radius: 3px;
  207. }
  208. /* 收缩为 3 行预览 */
  209. .clamp-5 {
  210. display: -webkit-box;
  211. -webkit-line-clamp: 3;
  212. -webkit-box-orient: vertical;
  213. overflow: hidden;
  214. }
  215. .message-box {
  216. margin-bottom: 0.25rem;
  217. }
  218. /* 让按钮文字不换行(包含导出按钮) */
  219. .stButton > button {
  220. white-space: nowrap;
  221. }
  222. </style>
  223. """, unsafe_allow_html=True)
  224. # 侧边栏:当前对话搜索(仅高亮,不提供“上/下”和“清除定位”按钮)
  225. def message_contains_query(content, query: str) -> bool:
  226. if not query:
  227. return False
  228. q = query.lower()
  229. if isinstance(content, str):
  230. return q in content.lower()
  231. elif isinstance(content, list):
  232. for part in content:
  233. if part.get("type") == "text" and q in part.get("text", "").lower():
  234. return True
  235. return False
  236. else:
  237. return q in str(content).lower()
  238. search_query = st.sidebar.text_input("搜索当前对话内容")
  239. # 统计匹配数量(仅展示数量,不提供跳转)
  240. if search_query:
  241. results = [i for i, m in enumerate(st.session_state.messages) if message_contains_query(m.get("content"), search_query)]
  242. st.sidebar.write(f"共找到 {len(results)} 条匹配。")
  243. else:
  244. st.sidebar.write("无匹配。")
  245. # 历史消息数量控制
  246. if "history_count" not in st.session_state:
  247. st.session_state.history_count = 0
  248. max_history_count = len(st.session_state.messages) if st.session_state.messages else 10
  249. st.sidebar.slider(
  250. "选择使用的历史消息数量(共" + str(len(st.session_state.messages)) + "条)",
  251. min_value=0,
  252. max_value=max_history_count,
  253. value=st.session_state.history_count,
  254. key="history_count"
  255. )
  256. st.sidebar.write(f"您选择的历史消息数量是: {st.session_state.history_count}")
  257. # 文件上传保存为 HTTP 路径
  258. BASE_URL = "http://117.50.195.224:8007/download/"
  259. def save_uploaded_file(uploaded_file):
  260. unique_filename = f"{uuid.uuid4()}_{uploaded_file.name}"
  261. file_path = os.path.join(upload_dir, unique_filename)
  262. with open(file_path, "wb") as f:
  263. f.write(uploaded_file.getvalue())
  264. return f"{BASE_URL}{unique_filename}"
  265. # ========== 高亮辅助函数 ==========
  266. def highlight_text_safe(text: str, query: str) -> str:
  267. if not query:
  268. return html.escape(text)
  269. pattern = re.compile(re.escape(query), re.IGNORECASE)
  270. parts = pattern.split(text)
  271. matches = pattern.findall(text)
  272. out = []
  273. for i, part in enumerate(parts):
  274. out.append(html.escape(part))
  275. if i < len(matches):
  276. out.append(f"<mark class='hl'>{html.escape(matches[i])}</mark>")
  277. return "".join(out)
  278. def render_content(content, query: str = None, clamp_lines: int = None):
  279. # clamp_lines 为 None 时不收缩;为整数时收缩为该行数(当前实现为 5)
  280. clamp_class = "clamp-5" if clamp_lines and clamp_lines > 0 else None
  281. def render_text(txt: str):
  282. if query:
  283. html_text = highlight_text_safe(txt, query)
  284. if clamp_class:
  285. st.markdown(f"<div class='message-box {clamp_class}'>{html_text}</div>", unsafe_allow_html=True)
  286. else:
  287. st.markdown(html_text, unsafe_allow_html=True)
  288. else:
  289. if clamp_class:
  290. st.markdown(f"<div class='message-box {clamp_class}'>{html.escape(txt)}</div>", unsafe_allow_html=True)
  291. else:
  292. st.markdown(txt)
  293. if isinstance(content, str):
  294. render_text(content)
  295. elif isinstance(content, list):
  296. for part in content:
  297. if part.get("type") == "text":
  298. txt = part.get("text", "")
  299. render_text(txt)
  300. elif part.get("type") == "image_url":
  301. url = part["image_url"]["url"]
  302. # 图片不收缩
  303. if "base64," in url:
  304. base64_data = url.split("base64,")[1]
  305. image_bytes = base64.b64decode(base64_data)
  306. st.image(image_bytes, caption="上传的图片")
  307. else:
  308. st.image(url, caption="上传的图片")
  309. else:
  310. # 其他类型直接显示
  311. if clamp_class:
  312. st.markdown(f"<div class='message-box {clamp_class}'>{html.escape(str(content))}</div>", unsafe_allow_html=True)
  313. else:
  314. st.markdown(str(content))
  315. else:
  316. if clamp_class:
  317. st.markdown(f"<div class='message-box {clamp_class}'>{html.escape(str(content))}</div>", unsafe_allow_html=True)
  318. else:
  319. st.markdown(str(content))
  320. # ========== 渲染完整聊天消息(支持高亮与“历史自动收缩”/“搜索时展开”) ==========
  321. messages_len = len(st.session_state.messages)
  322. for idx, message in enumerate(st.session_state.messages):
  323. # 搜索时全部展开;非搜索时,除了最后一条,其余收缩为 5 行
  324. is_searching = bool(search_query)
  325. default_clamp = None if is_searching else (5 if idx < messages_len - 1 else None)
  326. is_expanded = (idx in st.session_state.expanded_messages)
  327. clamp_lines = None if (is_searching or is_expanded) else default_clamp
  328. with st.chat_message(message["role"]):
  329. add_hit = bool(search_query and message_contains_query(message.get("content"), search_query))
  330. render_content(
  331. message.get("content"),
  332. query=search_query if add_hit else None,
  333. clamp_lines=clamp_lines
  334. )
  335. # 同一行显示“展开/收起”和“导出”按钮,并保证导出按钮文本不换行
  336. spacer_col, export_col, expand_col = st.columns([8, 1, 1])
  337. with expand_col:
  338. if default_clamp and not is_searching:
  339. if not is_expanded:
  340. if st.button("\>\>", key=f"expand_{idx}"):
  341. st.session_state.expanded_messages.add(idx)
  342. st.rerun()
  343. else:
  344. if st.button("<<", key=f"collapse_{idx}"):
  345. if idx in st.session_state.expanded_messages:
  346. st.session_state.expanded_messages.remove(idx)
  347. st.rerun()
  348. with export_col:
  349. if message["role"] == "assistant":
  350. export_key = f"export_{idx}"
  351. if st.button("导出", key=export_key):
  352. export_message(message.get("content"))
  353. # 注:根据需求“去掉上,下和清除定位按钮,只高亮即可”,不再自动滚动到定位消息,也不提供定位跳转控件。
  354. # ========== 处理 chat_input ==========
  355. prompt_input = st.chat_input(
  356. "What is up?",
  357. accept_file="multiple",
  358. file_type=["jpg", "jpeg", "png", "txt", "pdf", "doc", "docx"]
  359. )
  360. if prompt_input:
  361. text_content = prompt_input.text or ""
  362. uploaded_files = prompt_input.files
  363. content = text_content
  364. additional_prompt = ""
  365. if uploaded_files:
  366. content_parts = [{"type": "text", "text": text_content}]
  367. for file in uploaded_files:
  368. file_type = (file.type.split('/')[1].lower() if '/' in file.type else file.type.lower()).strip()
  369. if file_type in ["jpg", "jpeg", "png"]:
  370. image_bytes = file.getvalue()
  371. base64_image = base64.b64encode(image_bytes).decode('utf-8')
  372. content_parts.append({
  373. "type": "image_url",
  374. "image_url": {"url": f"data:image/{file_type};base64,{base64_image}"}
  375. })
  376. else:
  377. http_url = save_uploaded_file(file)
  378. additional_prompt += f"本次提问包含:{http_url} 文件\n"
  379. content = content_parts if len(content_parts) > 1 else text_content
  380. if additional_prompt:
  381. if isinstance(content, list):
  382. content[0]["text"] = (content[0]["text"] + "\n" + additional_prompt).strip()
  383. else:
  384. content = (content + "\n" + additional_prompt).strip()
  385. # 用户发送新消息后,“历史的自动收缩”:清空展开状态集合
  386. st.session_state.expanded_messages = set()
  387. st.session_state.messages.append({"role": "user", "content": content})
  388. with st.chat_message("user"):
  389. # 最新消息不收缩
  390. render_content(content, clamp_lines=None)
  391. with st.chat_message("assistant"):
  392. client.api_key = model_keys[selected_model]
  393. stream = output_mode == "流式输出 (Stream)"
  394. try:
  395. if st.session_state.history_count > 0:
  396. messages_to_send = [
  397. {"role": m["role"], "content": m["content"]}
  398. for m in st.session_state.messages[-st.session_state.history_count:]
  399. ]
  400. else:
  401. messages_to_send = [{"role": "user", "content": content}]
  402. res = client.chat.completions.create(
  403. model=selected_model,
  404. messages=messages_to_send,
  405. stream=stream,
  406. )
  407. if stream:
  408. assistant_message = st.write_stream(res)
  409. if assistant_message:
  410. st.session_state.messages.append({"role": "assistant", "content": assistant_message})
  411. else:
  412. st.warning("收到空响应。")
  413. else:
  414. if len(res.choices) > 0:
  415. assistant_message = res.choices[0].message.content
  416. if assistant_message:
  417. render_content(assistant_message)
  418. st.session_state.messages.append({"role": "assistant", "content": assistant_message})
  419. else:
  420. st.warning("收到空响应。")
  421. else:
  422. st.warning("响应格式不正确,未找到有效消息。")
  423. save_history()
  424. # 确保新增的最后一条消息在有搜索时也能被高亮与计数,强制重渲染
  425. st.rerun()
  426. except Exception as e:
  427. st.error(f"发生错误:{e}")
  428. # ========== 侧边栏操作 ==========
  429. st.sidebar.header("操作")
  430. if st.sidebar.button("New Chat"):
  431. st.session_state.messages = []
  432. st.session_state.session_id = load_session_id()
  433. st.session_state.session_id += 1
  434. save_session_id(st.session_state.session_id)
  435. # 新建会话时重置展开状态
  436. st.session_state.expanded_messages = set()
  437. st.success("当前会话已清空。")
  438. st.sidebar.header("历史对话")
  439. show_history_files(st.session_state.current_page)