# -*- coding: utf-8 -*- import os import re import glob import json import shutil import base64 import uuid import datetime import subprocess import html import streamlit as st import streamlit.components.v1 as components from openai import OpenAI # ========== 启动文件服务(非阻塞)========== try: if "fastpdf_spawned" not in st.session_state: subprocess.Popen(["python", "fastpdf.py"]) st.session_state["fastpdf_spawned"] = True except Exception: pass # ========== 示例:模型和相应的 API 密钥 ========== # 请在生产环境中改为安全地注入 API Key(如通过环境变量) default_key = "sk-re2NlaKIQn11ZNWzAbB6339cEbF94c6aAfC8B7Ab82879bEa" # 默认令牌,请勿用于生产 model_keys = { "grok-3": default_key, "grok-4.1": default_key, "grok-4.2": default_key, "gpt-5.4-mini": default_key, "gpt-4o-mini": default_key, # "gpt-4.1-mini-2025-04-14": default_key, "o1-mini": default_key, "o4-mini": default_key, "deepseek-v3.1": default_key, "deepseek-r1": default_key, "grok-4": default_key, "gpt-5-all": default_key, "gpt-4o-all": default_key, # "gpt-5-mini-2025-08-07": default_key, "o3-mini-all": default_key, # "claude-sonnet-4-20250514": default_key, } st.title("ChatGPT-like Clone") selected_model = st.sidebar.selectbox("选择模型", list(model_keys.keys())) api_key = model_keys[selected_model] api_url = "https://yunwu.ai/v1" client = OpenAI(api_key=api_key, base_url=api_url) data_dir = "data" backup_dir = "data_bak" blog_dir = "blog" upload_dir = "uploads" for directory in [data_dir, backup_dir, blog_dir, upload_dir]: if not os.path.exists(directory): os.makedirs(directory) session_id_file = os.path.join(data_dir, "session_id.txt") def load_session_id(): if os.path.exists(session_id_file): try: with open(session_id_file, "r", encoding="utf-8") as f: c = f.read().strip() return int(c) if c.isdigit() else 0 except Exception: return 0 return 0 def save_session_id(session_id: int): with open(session_id_file, "w", encoding="utf-8") as f: f.write(str(session_id)) if "openai_model" not in st.session_state: st.session_state["openai_model"] = selected_model def load_history(file_path: str): match = re.search(r'chat_history_(\d+)\.json', os.path.basename(file_path)) if match: st.session_state.session_id = int(match.group(1)) try: with open(file_path, "r", encoding="utf-8") as f: st.session_state.messages = json.load(f) except Exception: st.session_state.messages = [] def load_latest_history(): history_files = sorted(glob.glob(os.path.join(data_dir, "*.json")), key=os.path.getmtime) if history_files: load_history(history_files[-1]) if "messages" not in st.session_state: st.session_state.messages = [] load_latest_history() if "session_id" not in st.session_state: st.session_state.session_id = load_session_id() # 用于控制“单块展开”的状态集合(索引集合) if "expanded_messages" not in st.session_state: st.session_state.expanded_messages = set() def save_history(): file_path = os.path.join(data_dir, f"chat_history_{st.session_state.session_id}.json") try: json.dumps(st.session_state.messages) with open(file_path, "w", encoding="utf-8") as f: json.dump(st.session_state.messages, f, ensure_ascii=False) except (TypeError, ValueError): st.error("数据格式不正确,无法保存记录") def move_history(file_path: str): if os.path.exists(file_path): destination_file = os.path.join(backup_dir, os.path.basename(file_path)) shutil.move(file_path, destination_file) def delete_history(file_path: str): if os.path.exists(file_path): os.remove(file_path) def text_from_content(content): if isinstance(content, str): return content elif isinstance(content, list): return " ".join([part.get("text", "") for part in content if part.get("type") == "text"]) return str(content) def get_chat_title(messages): if messages: title = text_from_content(messages[0].get("content", "")) or "空的聊天" return title[:12] return "空的聊天"[:12] def show_history_files(page: int = 0, page_size: int = 10): with st.sidebar.expander("历史聊天记录"): history_files = [(f, os.path.getmtime(f)) for f in glob.glob(os.path.join(data_dir, "*.json"))] history_files.sort(key=lambda x: x[1], reverse=True) total_files = len(history_files) total_pages = (total_files // page_size) + (1 if total_files % page_size > 0 else 0) # 页码范围校正 if page >= total_pages and total_pages > 0: st.session_state.current_page = max(0, total_pages - 1) st.rerun() elif page < 0: st.session_state.current_page = 0 st.rerun() start = page * page_size end = start + page_size display_files = history_files[start:end] if not display_files: st.write("无记录。") else: for file_path, _mtime in display_files: try: with open(file_path, "r", encoding="utf-8") as f: chat_history = json.load(f) except Exception: continue file_name = get_chat_title(chat_history) col1, col2, col3 = st.sidebar.columns([3, 1, 1]) with col1: if st.button(file_name, key=f"load_{os.path.basename(file_path)}"): load_history(file_path) # 切换历史时,清空单块展开状态 st.session_state.expanded_messages = set() st.rerun() with col2: if st.button("📦", key=f"move_{os.path.basename(file_path)}", help="移动到备份文件夹"): move_history(file_path) st.success(f"{file_name} 已移动到备份。") st.rerun() with col3: if st.button("❌", key=f"delete_{os.path.basename(file_path)}", help="删除"): delete_history(file_path) st.success(f"{file_name} 已删除。") st.rerun() col_prev, col_next = st.sidebar.columns([1, 1]) with col_prev: if st.button("上一页", disabled=page <= 0, key="history_prev_page"): st.session_state.current_page -= 1 st.rerun() with col_next: if st.button("下一页", disabled=page >= total_pages - 1, key="history_next_page"): st.session_state.current_page += 1 st.rerun() if "current_page" not in st.session_state: st.session_state.current_page = 0 def export_message(content): # 文本提取与转义 if isinstance(content, str): processed_content = content.replace('\n', '\n').replace('"', '"') elif isinstance(content, list): processed_content = " ".join([part.get("text", "") for part in content if part.get("type") == "text"]) processed_content = processed_content.replace('\n', '\n').replace('"', '"') else: processed_content = str(content) timestamp = datetime.datetime.now().strftime("%m%d%H%M") first_10 = ( processed_content[:10] .replace(' ', '') .replace('/', '') .replace('\\', '') .replace(':', '') .replace('```', '') ) filename = f"{timestamp}_{first_10}.txt" file_path = os.path.join(blog_dir, filename) try: with open(file_path, "w", encoding="utf-8") as f: f.write(processed_content) st.success(f"已导出到 {file_path}") except Exception as e: st.error(f"导出失败:{e}") # 侧边栏:输出模式 output_mode = st.sidebar.selectbox("选择输出模式", ["流式输出 (Stream)", "非流式输出 (Non-stream)"]) # ========== 高亮样式与收缩样式 ========== st.markdown(""" """, unsafe_allow_html=True) # 侧边栏:当前对话搜索(仅高亮,不提供“上/下”和“清除定位”按钮) def message_contains_query(content, query: str) -> bool: if not query: return False q = query.lower() if isinstance(content, str): return q in content.lower() elif isinstance(content, list): for part in content: if part.get("type") == "text" and q in part.get("text", "").lower(): return True return False else: return q in str(content).lower() search_query = st.sidebar.text_input("搜索当前对话内容") # 统计匹配数量(仅展示数量,不提供跳转) if search_query: results = [i for i, m in enumerate(st.session_state.messages) if message_contains_query(m.get("content"), search_query)] st.sidebar.write(f"共找到 {len(results)} 条匹配。") else: st.sidebar.write("无匹配。") # 历史消息数量控制 if "history_count" not in st.session_state: st.session_state.history_count = 0 max_history_count = len(st.session_state.messages) if st.session_state.messages else 10 st.sidebar.slider( "选择使用的历史消息数量(共" + str(len(st.session_state.messages)) + "条)", min_value=0, max_value=max_history_count, value=st.session_state.history_count, key="history_count" ) st.sidebar.write(f"您选择的历史消息数量是: {st.session_state.history_count}") # 文件上传保存为 HTTP 路径 BASE_URL = "http://117.50.195.224:8007/download/" def save_uploaded_file(uploaded_file): unique_filename = f"{uuid.uuid4()}_{uploaded_file.name}" file_path = os.path.join(upload_dir, unique_filename) with open(file_path, "wb") as f: f.write(uploaded_file.getvalue()) return f"{BASE_URL}{unique_filename}" # ========== 高亮辅助函数 ========== def highlight_text_safe(text: str, query: str) -> str: if not query: return html.escape(text) pattern = re.compile(re.escape(query), re.IGNORECASE) parts = pattern.split(text) matches = pattern.findall(text) out = [] for i, part in enumerate(parts): out.append(html.escape(part)) if i < len(matches): out.append(f"{html.escape(matches[i])}") return "".join(out) def render_content(content, query: str = None, clamp_lines: int = None): # clamp_lines 为 None 时不收缩;为整数时收缩为该行数(当前实现为 5) clamp_class = "clamp-5" if clamp_lines and clamp_lines > 0 else None def render_text(txt: str): if query: html_text = highlight_text_safe(txt, query) if clamp_class: st.markdown(f"
", unsafe_allow_html=True) else: st.markdown(html_text, unsafe_allow_html=True) else: if clamp_class: st.markdown(f"", unsafe_allow_html=True) else: st.markdown(txt) if isinstance(content, str): render_text(content) elif isinstance(content, list): for part in content: if part.get("type") == "text": txt = part.get("text", "") render_text(txt) elif part.get("type") == "image_url": url = part["image_url"]["url"] # 图片不收缩 if "base64," in url: base64_data = url.split("base64,")[1] image_bytes = base64.b64decode(base64_data) st.image(image_bytes, caption="上传的图片") else: st.image(url, caption="上传的图片") else: # 其他类型直接显示 if clamp_class: st.markdown(f"", unsafe_allow_html=True) else: st.markdown(str(content)) else: if clamp_class: st.markdown(f"", unsafe_allow_html=True) else: st.markdown(str(content)) # ========== 渲染完整聊天消息(支持高亮与“历史自动收缩”/“搜索时展开”) ========== messages_len = len(st.session_state.messages) for idx, message in enumerate(st.session_state.messages): # 搜索时全部展开;非搜索时,除了最后一条,其余收缩为 5 行 is_searching = bool(search_query) default_clamp = None if is_searching else (5 if idx < messages_len - 1 else None) is_expanded = (idx in st.session_state.expanded_messages) clamp_lines = None if (is_searching or is_expanded) else default_clamp with st.chat_message(message["role"]): add_hit = bool(search_query and message_contains_query(message.get("content"), search_query)) render_content( message.get("content"), query=search_query if add_hit else None, clamp_lines=clamp_lines ) # 同一行显示“展开/收起”和“导出”按钮,并保证导出按钮文本不换行 spacer_col, export_col, expand_col = st.columns([8, 1, 1]) with expand_col: if default_clamp and not is_searching: if not is_expanded: if st.button("\>\>", key=f"expand_{idx}"): st.session_state.expanded_messages.add(idx) st.rerun() else: if st.button("<<", key=f"collapse_{idx}"): if idx in st.session_state.expanded_messages: st.session_state.expanded_messages.remove(idx) st.rerun() with export_col: if message["role"] == "assistant": export_key = f"export_{idx}" if st.button("导出", key=export_key): export_message(message.get("content")) # 注:根据需求“去掉上,下和清除定位按钮,只高亮即可”,不再自动滚动到定位消息,也不提供定位跳转控件。 # ========== 处理 chat_input ========== prompt_input = st.chat_input( "What is up?", accept_file="multiple", file_type=["jpg", "jpeg", "png", "txt", "pdf", "doc", "docx"] ) if prompt_input: text_content = prompt_input.text or "" uploaded_files = prompt_input.files content = text_content additional_prompt = "" if uploaded_files: content_parts = [{"type": "text", "text": text_content}] for file in uploaded_files: file_type = (file.type.split('/')[1].lower() if '/' in file.type else file.type.lower()).strip() if file_type in ["jpg", "jpeg", "png"]: image_bytes = file.getvalue() base64_image = base64.b64encode(image_bytes).decode('utf-8') content_parts.append({ "type": "image_url", "image_url": {"url": f"data:image/{file_type};base64,{base64_image}"} }) else: http_url = save_uploaded_file(file) additional_prompt += f"本次提问包含:{http_url} 文件\n" content = content_parts if len(content_parts) > 1 else text_content if additional_prompt: if isinstance(content, list): content[0]["text"] = (content[0]["text"] + "\n" + additional_prompt).strip() else: content = (content + "\n" + additional_prompt).strip() # 用户发送新消息后,“历史的自动收缩”:清空展开状态集合 st.session_state.expanded_messages = set() st.session_state.messages.append({"role": "user", "content": content}) with st.chat_message("user"): # 最新消息不收缩 render_content(content, clamp_lines=None) with st.chat_message("assistant"): client.api_key = model_keys[selected_model] stream = output_mode == "流式输出 (Stream)" try: if st.session_state.history_count > 0: messages_to_send = [ {"role": m["role"], "content": m["content"]} for m in st.session_state.messages[-st.session_state.history_count:] ] else: messages_to_send = [{"role": "user", "content": content}] res = client.chat.completions.create( model=selected_model, messages=messages_to_send, stream=stream, ) if stream: assistant_message = st.write_stream(res) if assistant_message: st.session_state.messages.append({"role": "assistant", "content": assistant_message}) else: st.warning("收到空响应。") else: if len(res.choices) > 0: assistant_message = res.choices[0].message.content if assistant_message: render_content(assistant_message) st.session_state.messages.append({"role": "assistant", "content": assistant_message}) else: st.warning("收到空响应。") else: st.warning("响应格式不正确,未找到有效消息。") save_history() # 确保新增的最后一条消息在有搜索时也能被高亮与计数,强制重渲染 st.rerun() except Exception as e: st.error(f"发生错误:{e}") # ========== 侧边栏操作 ========== st.sidebar.header("操作") if st.sidebar.button("New Chat"): st.session_state.messages = [] st.session_state.session_id = load_session_id() st.session_state.session_id += 1 save_session_id(st.session_state.session_id) # 新建会话时重置展开状态 st.session_state.expanded_messages = set() st.success("当前会话已清空。") st.sidebar.header("历史对话") show_history_files(st.session_state.current_page)