Skip to main content

codex_sync/
cc_switch.rs

1use anyhow::{Context, Result, bail};
2use rusqlite::{Connection, params};
3use serde::Serialize;
4use serde_json::Value;
5use std::{
6    collections::BTreeMap,
7    fs,
8    path::{Path, PathBuf},
9};
10use walkdir::WalkDir;
11
12#[derive(Debug)]
13pub struct Paths {
14    codex_home: PathBuf,
15    cc_switch_home: PathBuf,
16}
17
18#[derive(Debug, Serialize)]
19pub struct Report {
20    pub current_cc_switch_provider: String,
21    pub target_model_provider: String,
22    pub rollout_counts: BTreeMap<String, usize>,
23    pub database_counts: BTreeMap<String, usize>,
24    pub rollout_files_to_change: usize,
25    pub database_rows_to_change: usize,
26    #[serde(skip_serializing_if = "Option::is_none")]
27    pub backup: Option<PathBuf>,
28}
29
30impl Paths {
31    pub fn resolve(codex_home: Option<PathBuf>, cc_switch_home: Option<PathBuf>) -> Result<Self> {
32        let home = dirs::home_dir().context("无法确定用户主目录")?;
33        let paths = Self {
34            codex_home: codex_home.unwrap_or_else(|| home.join(".codex")),
35            cc_switch_home: cc_switch_home.unwrap_or_else(|| home.join(".cc-switch")),
36        };
37        if !paths.codex_home.is_dir() {
38            bail!("Codex 数据目录不存在:{}", paths.codex_home.display());
39        }
40        if !paths.cc_switch_home.join("cc-switch.db").is_file() {
41            bail!("未检测到 CC-Switch:{}", paths.cc_switch_home.display());
42        }
43        Ok(paths)
44    }
45}
46
47pub fn inspect(paths: &Paths) -> Result<Report> {
48    let current_cc_switch_provider = current_cc_switch_provider(paths)?;
49    let target_model_provider = active_model_provider(paths)?;
50    let (rollout_counts, rollout_files_to_change) = rollout_counts(paths, &target_model_provider)?;
51    let (database_counts, database_rows_to_change) =
52        database_counts(paths, &target_model_provider)?;
53    Ok(Report {
54        current_cc_switch_provider,
55        target_model_provider,
56        rollout_counts,
57        database_counts,
58        rollout_files_to_change,
59        database_rows_to_change,
60        backup: None,
61    })
62}
63
64pub fn merge(paths: &Paths, apply: bool) -> Result<Report> {
65    let mut report = inspect(paths)?;
66    if !apply || (report.rollout_files_to_change == 0 && report.database_rows_to_change == 0) {
67        return Ok(report);
68    }
69    let backup = backup_dir(paths)?;
70    fs::create_dir_all(&backup)?;
71    rewrite_rollouts(paths, &report.target_model_provider, &backup)?;
72    rewrite_database(paths, &report.target_model_provider, &backup)?;
73    report.backup = Some(backup);
74    let verified = inspect(paths)?;
75    if verified.rollout_files_to_change != 0 || verified.database_rows_to_change != 0 {
76        bail!("合并后验证失败;原始文件已备份,请勿继续使用并检查备份");
77    }
78    Ok(report)
79}
80
81pub fn print_report(report: &Report) {
82    println!("CC-Switch 当前账号:{}", report.current_cc_switch_provider);
83    println!("Codex 当前记录桶:{}", report.target_model_provider);
84    println!("会话文件分布:");
85    for (provider, count) in &report.rollout_counts {
86        println!("  {provider}: {count}");
87    }
88    println!("数据库记录分布:");
89    for (provider, count) in &report.database_counts {
90        println!("  {provider}: {count}");
91    }
92    println!(
93        "待合并:{} 个会话文件,{} 条数据库记录",
94        report.rollout_files_to_change, report.database_rows_to_change
95    );
96}
97
98fn active_model_provider(paths: &Paths) -> Result<String> {
99    let text = fs::read_to_string(paths.codex_home.join("config.toml"))
100        .context("无法读取 Codex config.toml")?;
101    let value: toml::Value = toml::from_str(&text).context("Codex config.toml 格式错误")?;
102    value
103        .get("model_provider")
104        .and_then(|v| v.as_str())
105        .map(str::to_owned)
106        .context("Codex config.toml 缺少 model_provider,无法确定当前记录桶")
107}
108
109fn current_cc_switch_provider(paths: &Paths) -> Result<String> {
110    let conn = Connection::open_with_flags(
111        paths.cc_switch_home.join("cc-switch.db"),
112        rusqlite::OpenFlags::SQLITE_OPEN_READ_ONLY,
113    )?;
114    conn.query_row(
115        "SELECT name FROM providers WHERE app_type='codex' AND is_current=1 LIMIT 1",
116        [],
117        |row| row.get(0),
118    )
119    .context("CC-Switch 没有当前 Codex 账号")
120}
121
122fn rollout_paths(paths: &Paths) -> Vec<PathBuf> {
123    ["sessions", "archived_sessions"]
124        .into_iter()
125        .flat_map(|dir| {
126            WalkDir::new(paths.codex_home.join(dir))
127                .follow_links(false)
128                .into_iter()
129                .filter_map(Result::ok)
130                .filter(|e| {
131                    e.file_type().is_file() && e.path().extension().is_some_and(|x| x == "jsonl")
132                })
133                .map(|e| e.into_path())
134                .collect::<Vec<_>>()
135        })
136        .collect()
137}
138
139fn rollout_counts(paths: &Paths, target: &str) -> Result<(BTreeMap<String, usize>, usize)> {
140    let mut counts = BTreeMap::new();
141    let mut changes = 0;
142    for path in rollout_paths(paths) {
143        if let Some(provider) = rollout_provider(&path)? {
144            *counts.entry(provider.clone()).or_insert(0) += 1;
145            if provider != target {
146                changes += 1;
147            }
148        }
149    }
150    Ok((counts, changes))
151}
152
153fn rollout_provider(path: &Path) -> Result<Option<String>> {
154    for line in fs::read_to_string(path)?.lines() {
155        let value: Value = serde_json::from_str(line)
156            .with_context(|| format!("JSONL 格式错误:{}", path.display()))?;
157        if value.get("type").and_then(Value::as_str) == Some("session_meta") {
158            return Ok(value
159                .pointer("/payload/model_provider")
160                .and_then(Value::as_str)
161                .map(str::to_owned));
162        }
163    }
164    Ok(None)
165}
166
167fn database_counts(paths: &Paths, target: &str) -> Result<(BTreeMap<String, usize>, usize)> {
168    let db = paths.codex_home.join("state_5.sqlite");
169    if !db.is_file() {
170        return Ok((BTreeMap::new(), 0));
171    }
172    let conn = Connection::open_with_flags(db, rusqlite::OpenFlags::SQLITE_OPEN_READ_ONLY)?;
173    let mut stmt =
174        conn.prepare("SELECT model_provider, COUNT(*) FROM threads GROUP BY model_provider")?;
175    let mut counts = BTreeMap::new();
176    let rows = stmt.query_map([], |row| {
177        Ok((row.get::<_, String>(0)?, row.get::<_, usize>(1)?))
178    })?;
179    for row in rows {
180        let (provider, count) = row?;
181        counts.insert(provider, count);
182    }
183    let changes = conn.query_row(
184        "SELECT COUNT(*) FROM threads WHERE model_provider <> ?1",
185        [target],
186        |row| row.get(0),
187    )?;
188    Ok((counts, changes))
189}
190
191fn backup_dir(paths: &Paths) -> Result<PathBuf> {
192    let timestamp = std::time::SystemTime::now()
193        .duration_since(std::time::UNIX_EPOCH)?
194        .as_secs();
195    Ok(paths
196        .codex_home
197        .join("codex-sync-backups")
198        .join(format!("cc-switch-merge-{timestamp}")))
199}
200
201fn rewrite_rollouts(paths: &Paths, target: &str, backup: &Path) -> Result<()> {
202    for path in rollout_paths(paths) {
203        if rollout_provider(&path)?
204            .as_deref()
205            .is_none_or(|p| p == target)
206        {
207            continue;
208        }
209        let rel = path.strip_prefix(&paths.codex_home)?;
210        let backup_path = backup.join(rel);
211        if let Some(parent) = backup_path.parent() {
212            fs::create_dir_all(parent)?;
213        }
214        fs::copy(&path, &backup_path)?;
215        let mut output = String::new();
216        for line in fs::read_to_string(&path)?.lines() {
217            let mut value: Value = serde_json::from_str(line)?;
218            if value.get("type").and_then(Value::as_str) == Some("session_meta")
219                && let Some(payload) = value.get_mut("payload").and_then(Value::as_object_mut)
220            {
221                payload.insert("model_provider".into(), Value::String(target.into()));
222            }
223            output.push_str(&serde_json::to_string(&value)?);
224            output.push('\n');
225        }
226        let temp = path.with_extension("jsonl.codex-sync-tmp");
227        fs::write(&temp, output)?;
228        fs::rename(temp, path)?;
229    }
230    Ok(())
231}
232
233fn rewrite_database(paths: &Paths, target: &str, backup: &Path) -> Result<()> {
234    let db = paths.codex_home.join("state_5.sqlite");
235    if !db.is_file() {
236        return Ok(());
237    }
238    let backup_db = backup.join("state_5.sqlite");
239    let mut conn = Connection::open(&db)?;
240    conn.execute(
241        "VACUUM INTO ?1",
242        params![backup_db.to_string_lossy().as_ref()],
243    )?;
244    let tx = conn.transaction()?;
245    tx.execute(
246        "UPDATE threads SET model_provider=?1 WHERE model_provider<>?1",
247        [target],
248    )?;
249    tx.commit()?;
250    Ok(())
251}
252
253#[cfg(test)]
254mod tests {
255    use super::*;
256
257    fn temp_dir() -> PathBuf {
258        std::env::temp_dir().join(format!(
259            "codex-sync-cc-switch-{}-{}",
260            std::process::id(),
261            std::time::SystemTime::now()
262                .duration_since(std::time::UNIX_EPOCH)
263                .unwrap()
264                .as_nanos()
265        ))
266    }
267
268    #[test]
269    fn missing_directories_are_rejected() {
270        assert!(
271            Paths::resolve(
272                Some(PathBuf::from("/missing-codex")),
273                Some(PathBuf::from("/missing-cc-switch"))
274            )
275            .is_err()
276        );
277    }
278
279    #[test]
280    fn merges_rollouts_and_database_with_backup() -> Result<()> {
281        let root = temp_dir();
282        let codex = root.join(".codex");
283        let cc = root.join(".cc-switch");
284        fs::create_dir_all(codex.join("sessions/2026/07/19"))?;
285        fs::create_dir_all(&cc)?;
286        fs::write(codex.join("config.toml"), "model_provider = \"current\"\n")?;
287        let rollout = codex.join("sessions/2026/07/19/rollout-test.jsonl");
288        fs::write(
289            &rollout,
290            "{\"type\":\"session_meta\",\"payload\":{\"id\":\"one\",\"model_provider\":\"old\"}}\n",
291        )?;
292
293        let state = Connection::open(codex.join("state_5.sqlite"))?;
294        state.execute(
295            "CREATE TABLE threads (id TEXT PRIMARY KEY, model_provider TEXT NOT NULL)",
296            [],
297        )?;
298        state.execute("INSERT INTO threads VALUES ('one', 'old')", [])?;
299        drop(state);
300        let switch = Connection::open(cc.join("cc-switch.db"))?;
301        switch.execute(
302            "CREATE TABLE providers (name TEXT, app_type TEXT, is_current INTEGER)",
303            [],
304        )?;
305        switch.execute(
306            "INSERT INTO providers VALUES ('Current Account', 'codex', 1)",
307            [],
308        )?;
309        drop(switch);
310
311        let paths = Paths {
312            codex_home: codex.clone(),
313            cc_switch_home: cc,
314        };
315        let preview = merge(&paths, false)?;
316        assert_eq!(preview.rollout_files_to_change, 1);
317        assert_eq!(preview.database_rows_to_change, 1);
318        let applied = merge(&paths, true)?;
319        assert!(applied.backup.as_ref().is_some_and(|p| p.is_dir()));
320        assert_eq!(rollout_provider(&rollout)?.as_deref(), Some("current"));
321        let state = Connection::open(codex.join("state_5.sqlite"))?;
322        let provider: String =
323            state.query_row("SELECT model_provider FROM threads", [], |r| r.get(0))?;
324        assert_eq!(provider, "current");
325        let verified = inspect(&paths)?;
326        assert_eq!(verified.rollout_files_to_change, 0);
327        assert_eq!(verified.database_rows_to_change, 0);
328        fs::remove_dir_all(root)?;
329        Ok(())
330    }
331}