commit 26265cbb1046f240a96a183e4c6fecb35ca4f93a Author: sunguosheng Date: Wed Jun 17 20:41:49 2026 +0800 Initial import of NLProg diff --git a/.gitattributes b/.gitattributes new file mode 100644 index 0000000..6313b56 --- /dev/null +++ b/.gitattributes @@ -0,0 +1 @@ +* text=auto eol=lf diff --git a/.gitignore b/.gitignore new file mode 100644 index 0000000..c63ee01 --- /dev/null +++ b/.gitignore @@ -0,0 +1,10 @@ +__pycache__/ +*.py[cod] +*$py.class +.pytest_cache/ +.mypy_cache/ +.venv/ +dist/ +build/ +*.egg-info/ +.nlprog/runs/*.json diff --git a/.nlprog/memory.md b/.nlprog/memory.md new file mode 100644 index 0000000..7f79e2d --- /dev/null +++ b/.nlprog/memory.md @@ -0,0 +1,6 @@ +# NLProg Memory + +- Add stable project facts here as the agent learns them. +- Keep secrets, tokens, and private credentials out of this file. + + diff --git a/.nlprog/project.json b/.nlprog/project.json new file mode 100644 index 0000000..11725c1 --- /dev/null +++ b/.nlprog/project.json @@ -0,0 +1,21 @@ +{ + "schema_version": 1, + "project_types": [ + "python" + ], + "important_files": [ + "pyproject.toml" + ], + "verification_commands": [ + "python -m compileall src", + "python -m unittest discover" + ], + "protected_paths": [ + ".git", + ".venv", + "node_modules", + "__pycache__", + "dist", + "build" + ] +} diff --git a/.nlprog/rules.md b/.nlprog/rules.md new file mode 100644 index 0000000..b2c1748 --- /dev/null +++ b/.nlprog/rules.md @@ -0,0 +1,22 @@ +# NLProg Rules + +## Editing + +- Inspect relevant files before editing. +- Prefer small, exact replacements over whole-file rewrites. +- Preview edits and ask for confirmation before writing files. +- Do not write secrets, API keys, or private credentials into the repository. + +## Protected Paths + +- `.git` +- `.venv` +- `node_modules` +- `__pycache__` +- `dist` +- `build` + +## Verification + +- `python -m compileall src` +- `python -m unittest discover` diff --git a/README.md b/README.md new file mode 100644 index 0000000..b79e79d --- /dev/null +++ b/README.md @@ -0,0 +1,234 @@ +# NLProg + +NLProg 是一个终端优先的自然语言编程 Agent。你可以用中文或英文描述任务,它会调用配置好的模型,检查项目、修改文件、运行验证命令,并在关键操作前请求确认。 + +## 功能 + +- 支持 `ask`、`chat`、`agent` 三种使用方式 +- 支持 OpenAI、OpenAI-compatible、Anthropic、Gemini、opencode、mock +- 可通过 opencode 使用本机已授权的 ChatGPT Pro/Plus +- Agent 支持文件查看、搜索、补丁式修改、命令执行 +- 写文件和高风险命令前会预览确认 +- 支持 JSON 自动修复重试 +- 支持项目规则、项目记忆、自动验证和运行日志 + +## 快速开始 + +```powershell +python -m pip install -e . +python -m nlprog init +python -m nlprog init-project --show +python -m nlprog doctor +python -m unittest discover +python -m nlprog agent "帮我检查这个项目" +``` + +Windows 一键安装: + +```powershell +.\install.ps1 +``` + +Windows 一键安装,并尝试安装 opencode: + +```powershell +.\install.ps1 -InstallOpencode +``` + +安装脚本会优先使用 Scoop 安装 opencode;如果没有 Scoop 但有 npm,则使用 `npm install -g opencode-ai`。如果两者都没有,会提示手动安装命令。 + +## 配置管理 + +查看配置: + +```powershell +python -m nlprog config show +``` + +修改配置: + +```powershell +python -m nlprog config set provider opencode +python -m nlprog config set model "" +python -m nlprog config set json_repair_retries 3 +``` + +快捷切换到 opencode: + +```powershell +python -m nlprog config use-opencode +``` + +指定 opencode 路径: + +```powershell +python -m nlprog config use-opencode --command "C:\Users\win\scoop\shims\opencode.exe" +``` + +## 模型注册表 + +查看已注册模型: +```powershell +python -m nlprog models list +``` + +添加并启用 opencode: +```powershell +python -m nlprog models add codex --provider opencode --use +``` + +添加 OpenAI-compatible 模型: +```powershell +python -m nlprog models add deepseek --provider openai-compatible --model deepseek-chat --base-url https://api.deepseek.com/v1 --api-key-env DEEPSEEK_API_KEY +``` + +切换模型: +```powershell +python -m nlprog models use deepseek +``` + +查看或删除模型: +```powershell +python -m nlprog models show deepseek +python -m nlprog models remove deepseek +``` + +## 使用 opencode + +先确认 opencode 已经登录: + +```powershell +opencode auth list +``` + +然后运行: + +```powershell +python -m nlprog config use-opencode +python -m nlprog agent "只列出项目根目录,不要修改文件" +``` + +## 项目规则和记忆 + +初始化项目配置: + +```powershell +python -m nlprog init-project --show +``` + +它会生成: + +```text +.nlprog/project.json +.nlprog/rules.md +.nlprog/memory.md +``` + +Agent 启动时会读取这些文件,用来了解项目类型、验证命令、保护目录、编辑规则和长期记忆。 + +## 自动验证和运行日志 + +Agent 修改文件后,会自动运行 `.nlprog/project.json` 里的验证命令,例如: + +```powershell +python -m compileall src +python -m unittest discover +``` + +跳过自动验证: + +```powershell +python -m nlprog agent "你的任务" --no-verify +``` + +每次 Agent 运行会保存 JSON 日志到: + +```text +.nlprog/runs/ +``` + +跳过日志: + +```powershell +python -m nlprog agent "你的任务" --no-log +``` + +## Patch 修改 + +Agent 可以用 `apply_patch` 做多文件修改。补丁会先预览,确认后才应用。 + +```text +*** Begin Patch +*** Add File: notes.txt ++hello +*** Update File: README.md +@@ + old line +-remove this ++add this +*** Delete File: obsolete.txt +*** End Patch +``` + +更新文件时,空格开头表示上下文,`-` 表示删除,`+` 表示新增。为了安全,旧代码块必须在文件中精确匹配一次。 + +## Doctor 诊断 + +检查本机配置、项目规则、验证命令和 opencode: + +```powershell +python -m nlprog doctor +``` + +跳过 opencode 检查: + +```powershell +python -m nlprog doctor --no-opencode +``` + +自动创建缺失的配置和项目文件: + +```powershell +python -m nlprog doctor --fix +``` + +## 运行日志 + +列出最近运行: + +```powershell +python -m nlprog runs list +``` + +查看最新日志: + +```powershell +python -m nlprog runs show latest +``` + +## 版本 + +```powershell +python -m nlprog version +``` + +## 命令安全策略 + +Agent 和 `ask` 执行命令前会检查风险: + +- 高风险系统命令会直接拦截,例如磁盘格式化、注册表修改、网络栈修改、关机重启 +- 删除文件、安装依赖、执行下载脚本等命令会要求额外确认 +- 普通验证命令,例如 `python -m compileall src`,会正常执行 + +## mock 模式 + +`mock` 不联网,适合验证程序流程: + +```powershell +python -m nlprog config set provider mock +python -m nlprog ask "列出当前目录" +``` + +## 安全说明 + +NLProg 会展示模型建议的操作,并在执行前要求确认。你仍然应该检查命令和补丁是否符合预期,尤其是删除文件、修改系统设置、上传数据、安装依赖这类操作。 diff --git a/install.ps1 b/install.ps1 new file mode 100644 index 0000000..403d23f --- /dev/null +++ b/install.ps1 @@ -0,0 +1,65 @@ +param( + [switch]$InstallOpencode +) + +$ErrorActionPreference = "Stop" + +function Test-Command($Name) { + return $null -ne (Get-Command $Name -ErrorAction SilentlyContinue) +} + +function Install-OpencodeIfRequested { + if (-not $InstallOpencode) { + Write-Host "Skipping opencode install. Use -InstallOpencode to install it." + return + } + + if (Test-Command "opencode") { + Write-Host "opencode is already installed." + opencode --version + return + } + + if (Test-Command "scoop") { + Write-Host "Installing opencode with Scoop..." + scoop install opencode + } + elseif (Test-Command "npm") { + Write-Host "Installing opencode with npm..." + npm install -g opencode-ai + } + else { + Write-Host "Could not find Scoop or npm." + Write-Host "Install one of them, then run:" + Write-Host " scoop install opencode" + Write-Host "or:" + Write-Host " npm install -g opencode-ai" + return + } + + if (Test-Command "opencode") { + Write-Host "opencode installed." + opencode --version + python -m nlprog config use-opencode + } + else { + Write-Host "opencode install command finished, but opencode is not on PATH yet." + Write-Host "Open a new terminal and run: opencode --version" + } +} + +Write-Host "Installing NLProg..." + +python --version +python -m pip install -e . + +python -m nlprog init +python -m nlprog init-project --show + +Install-OpencodeIfRequested + +python -m nlprog doctor + +Write-Host "" +Write-Host "NLProg is ready." +Write-Host "Try: python -m nlprog agent `"帮我检查这个项目`"" diff --git a/pyproject.toml b/pyproject.toml new file mode 100644 index 0000000..2ff0d9b --- /dev/null +++ b/pyproject.toml @@ -0,0 +1,20 @@ +[project] +name = "nlprog" +version = "0.1.0" +description = "A terminal-first natural language programming assistant with pluggable LLM providers." +readme = "README.md" +requires-python = ">=3.8" +dependencies = [] + +[project.scripts] +nlprog = "nlprog.cli:main" + +[build-system] +requires = ["setuptools>=68"] +build-backend = "setuptools.build_meta" + +[tool.setuptools.package-dir] +"" = "src" + +[tool.setuptools.packages.find] +where = ["src"] diff --git a/src/nlprog/__init__.py b/src/nlprog/__init__.py new file mode 100644 index 0000000..a05eb9a --- /dev/null +++ b/src/nlprog/__init__.py @@ -0,0 +1,3 @@ +__all__ = ["__version__"] + +__version__ = "0.1.0" diff --git a/src/nlprog/__main__.py b/src/nlprog/__main__.py new file mode 100644 index 0000000..2f05ddc --- /dev/null +++ b/src/nlprog/__main__.py @@ -0,0 +1,5 @@ +from .cli import main + + +if __name__ == "__main__": + main() diff --git a/src/nlprog/agent.py b/src/nlprog/agent.py new file mode 100644 index 0000000..8444bcc --- /dev/null +++ b/src/nlprog/agent.py @@ -0,0 +1,96 @@ +from __future__ import annotations + +import json +from dataclasses import dataclass, field +from pathlib import Path + +from .config import Config +from .context import collect_context +from .json_repair import complete_json, strip_fences +from .llm import LLMClient, Message + + +SYSTEM_PROMPT = """You are NLProg, a terminal-first natural language programming agent. +Return only JSON with this shape: +{ + "summary": "short explanation", + "commands": [{"cmd": "shell command", "reason": "why it is needed"}], + "notes": ["important caveats"] +} +Prefer small, reversible steps. Do not include destructive commands unless the user explicitly asked for them. +Use commands suitable for the user's current shell and operating system. +""" + + +@dataclass(frozen=True) +class CommandStep: + cmd: str + reason: str = "" + + +@dataclass(frozen=True) +class Plan: + summary: str + commands: list[CommandStep] = field(default_factory=list) + notes: list[str] = field(default_factory=list) + raw: str = "" + + +def build_plan(client: LLMClient, config: Config, root: Path, task: str) -> Plan: + context = collect_context(root, config.max_context_files) + user_prompt = f"""Current working directory: {root} + +Project context: +{context} + +User task: +{task} +""" + result = complete_json( + client, + [Message("system", SYSTEM_PROMPT), Message("user", user_prompt)], + schema_hint='{"summary": "short explanation", "commands": [{"cmd": "shell command", "reason": "why"}], "notes": ["caveat"]}', + validate=_validate_plan_json, + max_retries=config.json_repair_retries, + ) + return plan_from_json(result.data, result.raw) + + +def parse_plan(raw: str) -> Plan: + cleaned = strip_fences(raw.strip()) + try: + data = json.loads(cleaned) + except json.JSONDecodeError: + return Plan( + summary="The model did not return valid JSON. Review the raw response below.", + commands=[], + notes=[cleaned], + raw=raw, + ) + + return plan_from_json(data, raw) + + +def plan_from_json(data: dict[str, object], raw: str = "") -> Plan: + commands_raw = data.get("commands", []) + commands_list = commands_raw if isinstance(commands_raw, list) else [] + commands = [ + CommandStep(cmd=str(item.get("cmd", "")).strip(), reason=str(item.get("reason", "")).strip()) + for item in commands_list + if isinstance(item, dict) and str(item.get("cmd", "")).strip() + ] + notes_raw = data.get("notes", []) + notes = [str(note) for note in notes_raw] if isinstance(notes_raw, list) else [] + return Plan(summary=str(data.get("summary", "")), commands=commands, notes=notes, raw=raw) + + +def _validate_plan_json(data: dict[str, object]) -> str | None: + if "summary" not in data: + return "Missing required field: summary." + commands = data.get("commands") + if commands is not None and not isinstance(commands, list): + return "Field commands must be a list." + notes = data.get("notes") + if notes is not None and not isinstance(notes, list): + return "Field notes must be a list." + return None diff --git a/src/nlprog/agent_loop.py b/src/nlprog/agent_loop.py new file mode 100644 index 0000000..05deb62 --- /dev/null +++ b/src/nlprog/agent_loop.py @@ -0,0 +1,316 @@ +from __future__ import annotations + +from dataclasses import dataclass, field +from pathlib import Path +from typing import Any + +from .config import Config +from .json_repair import complete_json, parse_json_object, strip_fences +from .llm import LLMClient, Message +from .project_profile import append_memory, build_project_context, load_project_profile +from .tools import ToolBox + + +AGENT_SYSTEM_PROMPT = """You are NLProg Agent, a careful coding agent running in a terminal. +You can inspect and edit the current project by choosing one tool action per turn. + +Return only JSON. Use one of these shapes: +{"action": "list_files", "args": {"path": "."}, "reason": "why"} +{"action": "read_file", "args": {"path": "src/app.py"}, "reason": "why"} +{"action": "search_text", "args": {"query": "TODO", "path": "."}, "reason": "why"} +{"action": "create_file", "args": {"path": "file.txt", "content": "..."}, "reason": "why"} +{"action": "replace_in_file", "args": {"path": "file.txt", "old": "...", "new": "..."}, "reason": "why"} +{"action": "apply_patch", "args": {"patch": "*** Begin Patch\n*** Update File: path/to/file.py\n@@\n old line\n-new line to remove\n+new line to add\n*** End Patch"}, "reason": "why"} +{"action": "run_command", "args": {"cmd": "python -m pytest"}, "reason": "why"} +{"final": "what changed, verification performed, and any remaining caveats"} + +Rules: +- Inspect files before editing. +- Prefer apply_patch for code changes, especially multi-file edits. +- Use small exact replacements for tiny single-location edits. +- Run focused verification after edits when possible. +- Prefer suggested verification commands from the project context. +- Do not use destructive commands unless the user explicitly asked for them. +- If a tool result shows an error, adapt and continue. +""" + +AGENT_SCHEMA_HINT = """One of: +{"action": "list_files", "args": {"path": "."}, "reason": "why"} +{"action": "read_file", "args": {"path": "src/app.py"}, "reason": "why"} +{"action": "search_text", "args": {"query": "TODO", "path": "."}, "reason": "why"} +{"action": "create_file", "args": {"path": "file.txt", "content": "..."}, "reason": "why"} +{"action": "replace_in_file", "args": {"path": "file.txt", "old": "...", "new": "..."}, "reason": "why"} +{"action": "apply_patch", "args": {"patch": "*** Begin Patch\\n...\\n*** End Patch"}, "reason": "why"} +{"action": "run_command", "args": {"cmd": "python -m pytest"}, "reason": "why"} +{"final": "summary"} +""" + + +@dataclass(frozen=True) +class AgentEvent: + kind: str + message: str + + +@dataclass(frozen=True) +class AgentRun: + final: str + events: list[AgentEvent] = field(default_factory=list) + log_path: Path | None = None + + +def run_agent( + client: LLMClient, + config: Config, + root: Path, + task: str, + max_steps: int = 12, + confirm_edits: bool = True, + auto_approve_edits: bool = False, + update_memory: bool = True, + auto_verify: bool = True, + log_run: bool = True, +) -> AgentRun: + from .run_log import now_timestamp, write_run_log + + started_at = now_timestamp() + profile = load_project_profile(root) + toolbox = ToolBox( + root, + config.timeout_seconds, + protected_paths=profile.protected_paths, + auto_approve_commands=auto_approve_edits, + ) + project_context = build_project_context(root) + messages = [ + Message("system", AGENT_SYSTEM_PROMPT), + Message( + "user", + f"Current working directory: {root.resolve()}\n\nProject context:\n{project_context}\n\nTask:\n{task}", + ), + ] + events: list[AgentEvent] = [] + changed_since_verification = False + verification_failed = False + + for step in range(1, max_steps + 1): + json_result = complete_json( + client, + messages, + schema_hint=AGENT_SCHEMA_HINT, + validate=_validate_agent_json, + max_retries=config.json_repair_retries, + ) + data = json_result.data + if json_result.repaired: + events.append(AgentEvent("repair", f"JSON repaired after {json_result.attempts} attempts.")) + if "final" in data: + final = str(data.get("final", "")).strip() or "Done." + if auto_verify and changed_since_verification and profile.verification_commands: + verification_ok, verification_text = _run_verification(toolbox, profile.verification_commands) + events.append(AgentEvent("verification", verification_text)) + messages.append(Message("assistant", json_result.raw)) + messages.append( + Message( + "user", + f"Automatic verification result ({'ok' if verification_ok else 'error'}):\n{verification_text[:16000]}", + ) + ) + changed_since_verification = False + verification_failed = not verification_ok + if not verification_ok: + events.append(AgentEvent("action", "Verification failed; asking the model to fix it.")) + continue + + if verification_failed: + events.append(AgentEvent("action", "Waiting for a fix after failed verification.")) + messages.append(Message("assistant", json_result.raw)) + messages.append(Message("user", "Verification previously failed. Fix the issue or explain why it cannot be fixed.")) + verification_failed = False + continue + + events.append(AgentEvent("final", final)) + if update_memory: + _maybe_update_memory(client, root, task, final, messages, events, auto_approve_edits) + log_path = write_run_log(root, task, final, events, started_at) if log_run else None + if log_path is not None: + events.append(AgentEvent("log", f"Saved run log: {log_path.relative_to(root.resolve())}")) + return AgentRun(final=final, events=events, log_path=log_path) + + action = str(data.get("action", "")).strip() + args = data.get("args", {}) + reason = str(data.get("reason", "")).strip() + if not isinstance(args, dict): + args = {} + + events.append(AgentEvent("action", f"{step}. {action}: {reason}".strip())) + try: + pending_edit = toolbox.preview_edit(action, args) if confirm_edits else None + if pending_edit is not None: + events.append(AgentEvent("preview", pending_edit.preview)) + if not auto_approve_edits and not _confirm("Approve this action?"): + result_text = "Action was rejected by the user." + ok = False + else: + result = toolbox.run(action, args, approved=True) + result_text = result.output + ok = result.ok + else: + result = toolbox.run(action, args) + result_text = result.output + ok = result.ok + except Exception as exc: + result_text = f"Tool raised {exc.__class__.__name__}: {exc}" + ok = False + + status = "ok" if ok else "error" + events.append(AgentEvent(status, result_text)) + if ok and action in {"create_file", "replace_in_file", "apply_patch"}: + changed_since_verification = True + verification_failed = False + messages.append(Message("assistant", json_result.raw)) + messages.append(Message("user", f"Tool result ({status}) for {action}:\n{result_text[:16000]}")) + + final = f"Stopped after {max_steps} steps without a final answer." + events.append(AgentEvent("final", final)) + if update_memory: + _maybe_update_memory(client, root, task, final, messages, events, auto_approve_edits) + log_path = write_run_log(root, task, final, events, started_at) if log_run else None + if log_path is not None: + events.append(AgentEvent("log", f"Saved run log: {log_path.relative_to(root.resolve())}")) + return AgentRun(final=final, events=events, log_path=log_path) + + +def _confirm(prompt: str) -> bool: + answer = input(f"{prompt} [y/N] ").strip().lower() + return answer in {"y", "yes"} + + +def _run_verification(toolbox: ToolBox, commands: list[str]) -> tuple[bool, str]: + rows: list[str] = [] + all_ok = True + for command in commands: + result = toolbox.run("run_command", {"cmd": command}, approved=True) + rows.append(f"$ {command}") + rows.append(result.output) + if not result.ok: + all_ok = False + break + return all_ok, "\n".join(rows) + + +def _maybe_update_memory( + client: LLMClient, + root: Path, + task: str, + final: str, + messages: list[Message], + events: list[AgentEvent], + auto_approve: bool, +) -> None: + summary = _summarize_events(events) + prompt = f"""Decide whether this completed coding-agent run produced stable project memory. + +Return only JSON: +{{"items": ["short stable fact to remember"]}} + +Only include durable facts useful for future work, such as verification commands that worked, project structure, conventions, or recurring pitfalls. +Do not include secrets, credentials, tokens, one-off status, timestamps, vague praise, or anything speculative. +Return {{"items": []}} if there is nothing worth remembering. + +Task: +{task} + +Final answer: +{final} + +Run summary: +{summary} +""" + result = complete_json( + client, + [Message("system", "You extract safe long-term project memory."), Message("user", prompt)], + schema_hint='{"items": ["short stable fact to remember"]}', + validate=_validate_memory_json, + max_retries=2, + ) + data = result.data + if result.repaired: + events.append(AgentEvent("repair", f"Memory JSON repaired after {result.attempts} attempts.")) + items = _sanitize_memory_items(data.get("items", [])) + if not items: + events.append(AgentEvent("memory", "No new long-term memory suggested.")) + return + + preview = "\n".join(f"- {item}" for item in items) + events.append(AgentEvent("memory_preview", preview)) + if not auto_approve and not _confirm("Append these items to .nlprog/memory.md?"): + events.append(AgentEvent("memory", "Memory update skipped by the user.")) + return + + path = append_memory(root, items) + events.append(AgentEvent("memory", f"Updated {path.relative_to(root.resolve())}.")) + + +def _summarize_events(events: list[AgentEvent]) -> str: + rows: list[str] = [] + for event in events: + if event.kind in {"action", "ok", "error", "final"}: + rows.append(f"{event.kind}: {event.message[:1000]}") + return "\n".join(rows[-12:]) + + +def _sanitize_memory_items(value: object) -> list[str]: + if not isinstance(value, list): + return [] + items: list[str] = [] + for item in value: + text = " ".join(str(item).strip().split()) + if not text or len(text) > 240 or _looks_sensitive(text): + continue + if text not in items: + items.append(text) + return items[:8] + + +def _looks_sensitive(text: str) -> bool: + lowered = text.lower() + markers = ["api_key", "apikey", "token", "secret", "password", "bearer ", "sk-", "credential"] + return any(marker in lowered for marker in markers) + + +def _parse_json(raw: str) -> dict[str, Any]: + data, error = parse_json_object(raw) + if data is None: + return {"final": "The model did not return valid JSON. " + error + "\nRaw response:\n" + raw} + return data + + +def _strip_fences(text: str) -> str: + return strip_fences(text) + + +def _validate_agent_json(data: dict[str, Any]) -> str | None: + has_action = "action" in data + has_final = "final" in data + if has_action == has_final: + return "Return exactly one of action or final." + if has_final: + if not isinstance(data.get("final"), str): + return "Field final must be a string." + return None + if not isinstance(data.get("action"), str) or not data.get("action"): + return "Field action must be a non-empty string." + if "args" in data and not isinstance(data.get("args"), dict): + return "Field args must be an object." + if "reason" in data and not isinstance(data.get("reason"), str): + return "Field reason must be a string." + return None + + +def _validate_memory_json(data: dict[str, Any]) -> str | None: + if "items" not in data: + return "Missing required field: items." + if not isinstance(data.get("items"), list): + return "Field items must be a list." + return None diff --git a/src/nlprog/cli.py b/src/nlprog/cli.py new file mode 100644 index 0000000..7b1cafa --- /dev/null +++ b/src/nlprog/cli.py @@ -0,0 +1,480 @@ +from __future__ import annotations + +import argparse +import json +import sys +from pathlib import Path + +from .agent_loop import run_agent +from .agent import Plan, build_plan +from .command_safety import assess_command +from .config import ( + add_model, + list_models, + load_config, + read_config_file, + remove_model, + set_config_value, + use_model, + use_opencode_config, + write_default_config, +) +from .doctor import run_doctor +from .executor import run_command +from .llm import LLMError, create_client +from .project_profile import init_project, build_project_context +from .run_log import list_run_logs, load_run_log +from . import __version__ + + +def main(argv: list[str] | None = None) -> None: + _configure_stdio() + parser = argparse.ArgumentParser(prog="nlprog", description="Natural language programming in your terminal.") + subparsers = parser.add_subparsers(dest="command") + + subparsers.add_parser("init", help="Create a default config file.") + subparsers.add_parser("version", help="Show NLProg version.") + + config_parser = subparsers.add_parser("config", help="Show or update user configuration.") + config_subparsers = config_parser.add_subparsers(dest="config_command") + + config_subparsers.add_parser("show", help="Show the user config file.") + + config_set_parser = config_subparsers.add_parser("set", help="Set a config value.") + config_set_parser.add_argument("key", help="Config key.") + config_set_parser.add_argument("value", help="Config value.") + + use_opencode_parser = config_subparsers.add_parser("use-opencode", help="Use opencode as the default provider.") + use_opencode_parser.add_argument("--model", default="", help="Optional opencode model, such as openai/gpt-5-codex.") + use_opencode_parser.add_argument( + "--command", + dest="opencode_command", + default="opencode", + help="opencode executable path or command.", + ) + + models_parser = subparsers.add_parser("models", help="Manage named model profiles.") + models_subparsers = models_parser.add_subparsers(dest="models_command") + + models_subparsers.add_parser("list", help="List registered models.") + + models_show_parser = models_subparsers.add_parser("show", help="Show one registered model.") + models_show_parser.add_argument("name", help="Model profile name.") + + models_add_parser = models_subparsers.add_parser("add", help="Add or update a model profile.") + models_add_parser.add_argument("name", help="Model profile name.") + models_add_parser.add_argument("--provider", required=True, help="Provider: opencode, openai, anthropic, gemini, etc.") + models_add_parser.add_argument("--model", default="", help="Provider model id.") + models_add_parser.add_argument("--base-url", default=None, help="OpenAI-compatible base URL.") + models_add_parser.add_argument("--api-key-env", default=None, help="Environment variable containing the API key.") + models_add_parser.add_argument("--opencode-command", default=None, help="opencode executable path or command.") + models_add_parser.add_argument("--use", action="store_true", help="Activate this model after saving it.") + + models_use_parser = models_subparsers.add_parser("use", help="Activate a registered model.") + models_use_parser.add_argument("name", help="Model profile name.") + + models_remove_parser = models_subparsers.add_parser("remove", help="Remove a registered model.") + models_remove_parser.add_argument("name", help="Model profile name.") + + init_project_parser = subparsers.add_parser("init-project", help="Create editable project rules and memory.") + init_project_parser.add_argument("--cwd", default=".", help="Project directory.") + init_project_parser.add_argument("--force", action="store_true", help="Overwrite existing .nlprog files.") + init_project_parser.add_argument("--show", action="store_true", help="Print the detected project context.") + + doctor_parser = subparsers.add_parser("doctor", help="Check local NLProg setup.") + doctor_parser.add_argument("--cwd", default=".", help="Project directory.") + doctor_parser.add_argument("--no-opencode", action="store_true", help="Skip opencode checks.") + doctor_parser.add_argument("--fix", action="store_true", help="Create missing config and project profile files.") + + runs_parser = subparsers.add_parser("runs", help="List or show Agent run logs.") + runs_subparsers = runs_parser.add_subparsers(dest="runs_command") + + runs_list_parser = runs_subparsers.add_parser("list", help="List recent run logs.") + runs_list_parser.add_argument("--cwd", default=".", help="Project directory.") + runs_list_parser.add_argument("--limit", type=int, default=10, help="Number of logs to show.") + + runs_show_parser = runs_subparsers.add_parser("show", help="Show one run log.") + runs_show_parser.add_argument("run_id", nargs="?", default="latest", help="Run id, filename, prefix, or latest.") + runs_show_parser.add_argument("--cwd", default=".", help="Project directory.") + + ask_parser = subparsers.add_parser("ask", help="Ask NLProg to plan and optionally run commands.") + ask_parser.add_argument("task", nargs="+", help="The natural language task.") + ask_parser.add_argument("--cwd", default=".", help="Project directory.") + ask_parser.add_argument("--dry-run", action="store_true", help="Only print the plan.") + ask_parser.add_argument("--yes", action="store_true", help="Run commands without asking.") + + chat_parser = subparsers.add_parser("chat", help="Start an interactive session.") + chat_parser.add_argument("--cwd", default=".", help="Project directory.") + + agent_parser = subparsers.add_parser("agent", help="Run the tool-using coding agent.") + agent_parser.add_argument("task", nargs="+", help="The natural language task.") + agent_parser.add_argument("--cwd", default=".", help="Project directory.") + agent_parser.add_argument("--max-steps", type=int, default=12, help="Maximum tool-use steps.") + agent_parser.add_argument("--yes", action="store_true", help="Approve file edits without asking.") + agent_parser.add_argument("--no-confirm-edits", action="store_true", help="Disable edit previews and confirmations.") + agent_parser.add_argument("--no-memory", action="store_true", help="Skip automatic memory suggestions.") + agent_parser.add_argument("--no-verify", action="store_true", help="Skip automatic verification after edits.") + agent_parser.add_argument("--no-log", action="store_true", help="Do not save a run log.") + + args = parser.parse_args(argv) + + if args.command == "init": + path = write_default_config() + print(f"Config ready: {path}") + return + + if args.command == "version": + print(__version__) + return + + if args.command == "config": + _handle_config(args) + return + + if args.command == "models": + _handle_models(args) + return + + if args.command == "init-project": + cwd = Path(args.cwd).resolve() + created = init_project(cwd, force=args.force) + if created: + print("Project profile ready:") + for path in created: + print(f"- {path}") + else: + print("Project profile already exists. Use --force to overwrite.") + if args.show: + print("\nDetected project context:") + print(build_project_context(cwd)) + return + + if args.command == "doctor": + _handle_doctor(Path(args.cwd).resolve(), check_opencode=not args.no_opencode, fix=args.fix) + return + + if args.command == "runs": + _handle_runs(args) + return + + if args.command == "ask": + task = " ".join(args.task) + _handle_task(task, Path(args.cwd).resolve(), dry_run=args.dry_run, yes=args.yes) + return + + if args.command == "chat": + _chat(Path(args.cwd).resolve()) + return + + if args.command == "agent": + task = " ".join(args.task) + _handle_agent( + task, + Path(args.cwd).resolve(), + max_steps=args.max_steps, + confirm_edits=not args.no_confirm_edits, + auto_approve_edits=args.yes, + update_memory=not args.no_memory, + auto_verify=not args.no_verify, + log_run=not args.no_log, + ) + return + + parser.print_help() + + +def _chat(cwd: Path) -> None: + print("NLProg chat. Type /exit to quit.") + while True: + try: + task = input("\n> ").strip() + except (EOFError, KeyboardInterrupt): + print() + return + if not task: + continue + if task in {"/exit", "/quit"}: + return + _handle_task(task, cwd, dry_run=False, yes=False) + + +def _handle_task(task: str, cwd: Path, dry_run: bool, yes: bool) -> None: + config = load_config() + client = create_client(config) + try: + plan = build_plan(client, config, cwd, task) + except LLMError as exc: + print(f"LLM error: {exc}", file=sys.stderr) + raise SystemExit(1) from exc + + _print_plan(plan) + if dry_run or not plan.commands: + return + + should_confirm = config.require_confirmation and not yes + if should_confirm and not _confirm("Run these commands?"): + print("Cancelled.") + return + + for step in plan.commands: + assessment = assess_command(step.cmd) + if assessment.blocked: + print(f"\nBlocked: {step.cmd}", file=sys.stderr) + print(f"Reason: {assessment.reason}", file=sys.stderr) + break + if assessment.needs_confirmation and not yes: + print(f"\nCommand requires confirmation: {step.cmd}") + print(f"Reason: {assessment.reason}") + if not _confirm("Run this command?"): + print("Cancelled.") + break + print(f"\n$ {step.cmd}") + try: + result = run_command(step, cwd, config.timeout_seconds) + except Exception as exc: + print(f"Command failed to start: {exc}", file=sys.stderr) + break + if result.stdout: + print(result.stdout.rstrip()) + if result.stderr: + print(result.stderr.rstrip(), file=sys.stderr) + if result.returncode != 0: + print(f"Stopped because the command exited with code {result.returncode}.", file=sys.stderr) + break + + +def _handle_agent( + task: str, + cwd: Path, + max_steps: int, + confirm_edits: bool, + auto_approve_edits: bool, + update_memory: bool, + auto_verify: bool, + log_run: bool, +) -> None: + config = load_config() + client = create_client(config) + try: + run = run_agent( + client, + config, + cwd, + task, + max_steps=max_steps, + confirm_edits=confirm_edits, + auto_approve_edits=auto_approve_edits, + update_memory=update_memory, + auto_verify=auto_verify, + log_run=log_run, + ) + except LLMError as exc: + print(f"LLM error: {exc}", file=sys.stderr) + raise SystemExit(1) from exc + + for event in run.events: + if event.kind == "action": + print(f"\n> {event.message}") + elif event.kind == "preview": + print("\nAction preview:") + print(_indent(event.message, " ")) + elif event.kind == "ok": + print(_indent(event.message, " ")) + elif event.kind == "error": + print(_indent(event.message, " ERROR: "), file=sys.stderr) + elif event.kind == "final": + print(f"\nFinal: {event.message}") + elif event.kind == "repair": + print(_indent(event.message, " ")) + elif event.kind == "memory_preview": + print("\nMemory preview:") + print(_indent(event.message, " ")) + elif event.kind == "memory": + print(_indent(event.message, " ")) + elif event.kind == "verification": + print("\nVerification:") + print(_indent(event.message, " ")) + elif event.kind == "log": + print(_indent(event.message, " ")) + + +def _handle_doctor(cwd: Path, check_opencode: bool, fix: bool) -> None: + if fix: + config_path = write_default_config() + created = init_project(cwd) + print(f"Config ready: {config_path}") + if created: + print("Project files created:") + for path in created: + print(f"- {path}") + else: + print("Project profile already exists.") + + checks = run_doctor(cwd, check_opencode=check_opencode) + for check in checks: + print(f"[{check.status}] {check.name}: {check.detail}") + + failed = any(check.status == "FAIL" for check in checks) + warned = any(check.status == "WARN" for check in checks) + if failed: + raise SystemExit(2) + if warned: + raise SystemExit(1) + + +def _handle_runs(args: argparse.Namespace) -> None: + if args.runs_command == "list": + root = Path(args.cwd).resolve() + logs = list_run_logs(root) + if not logs: + print("No run logs found.") + return + for path in logs[: args.limit]: + try: + data = json.loads(path.read_text(encoding="utf-8")) + task = str(data.get("task", "")).replace("\n", " ") + ended_at = str(data.get("ended_at", "")) + except Exception: + task = "[could not read]" + ended_at = "" + print(f"{path.stem} {ended_at} {task[:100]}") + return + + if args.runs_command == "show": + root = Path(args.cwd).resolve() + try: + path, data = load_run_log(root, args.run_id) + except (FileNotFoundError, ValueError) as exc: + print(f"Run log error: {exc}", file=sys.stderr) + raise SystemExit(2) from exc + print(f"Run: {path.stem}") + print(f"Task: {data.get('task', '')}") + print(f"Started: {data.get('started_at', '')}") + print(f"Ended: {data.get('ended_at', '')}") + print(f"Final: {data.get('final', '')}") + print("\nEvents:") + for event in data.get("events", []): + print(f"- {event.get('kind', '')}: {str(event.get('message', '')).splitlines()[0][:160]}") + return + + print("Missing runs command. Use: list or show.", file=sys.stderr) + raise SystemExit(2) + + +def _handle_config(args: argparse.Namespace) -> None: + if args.config_command == "show": + path = write_default_config() + data = read_config_file() + print(f"Config: {path}") + print(json.dumps(data, indent=2, ensure_ascii=False)) + return + + if args.config_command == "set": + try: + path = set_config_value(args.key, args.value) + except ValueError as exc: + print(f"Config error: {exc}", file=sys.stderr) + raise SystemExit(2) from exc + print(f"Updated {args.key} in {path}") + return + + if args.config_command == "use-opencode": + path = use_opencode_config(model=args.model, command=args.opencode_command) + print(f"Configured opencode provider in {path}") + return + + print("Missing config command. Use: show, set, or use-opencode.", file=sys.stderr) + raise SystemExit(2) + + +def _handle_models(args: argparse.Namespace) -> None: + if args.models_command == "list": + active, models = list_models() + if not models: + print("No models registered.") + return + for name in sorted(models): + entry = models[name] + marker = "*" if name == active else " " + provider = entry.get("provider", "") + model = entry.get("model", "") + display_model = model if model else "(provider default)" + print(f"{marker} {name}: {provider} / {display_model}") + return + + if args.models_command == "show": + active, models = list_models() + if args.name not in models: + print(f"Model error: unknown model {args.name}", file=sys.stderr) + raise SystemExit(2) + marker = "active" if args.name == active else "inactive" + print(f"Model: {args.name} ({marker})") + print(json.dumps(models[args.name], indent=2, ensure_ascii=False)) + return + + if args.models_command == "add": + try: + path = add_model( + args.name, + provider=args.provider, + model=args.model, + base_url=args.base_url, + api_key_env=args.api_key_env, + opencode_command=args.opencode_command, + activate=args.use, + ) + except ValueError as exc: + print(f"Model error: {exc}", file=sys.stderr) + raise SystemExit(2) from exc + action = "Added and activated" if args.use else "Added" + print(f"{action} model {args.name} in {path}") + return + + if args.models_command == "use": + try: + path = use_model(args.name) + except ValueError as exc: + print(f"Model error: {exc}", file=sys.stderr) + raise SystemExit(2) from exc + print(f"Activated model {args.name} in {path}") + return + + if args.models_command == "remove": + try: + path = remove_model(args.name) + except ValueError as exc: + print(f"Model error: {exc}", file=sys.stderr) + raise SystemExit(2) from exc + print(f"Removed model {args.name} from {path}") + return + + print("Missing models command. Use: list, show, add, use, or remove.", file=sys.stderr) + raise SystemExit(2) + + +def _print_plan(plan: Plan) -> None: + print(f"\nSummary: {plan.summary or '(none)'}") + if plan.commands: + print("\nCommands:") + for index, step in enumerate(plan.commands, start=1): + reason = f" # {step.reason}" if step.reason else "" + print(f"{index}. {step.cmd}{reason}") + if plan.notes: + print("\nNotes:") + for note in plan.notes: + print(f"- {note}") + + +def _confirm(prompt: str) -> bool: + answer = input(f"{prompt} [y/N] ").strip().lower() + return answer in {"y", "yes"} + + +def _indent(text: str, prefix: str) -> str: + return "\n".join(prefix + line for line in text.splitlines()) + + +def _configure_stdio() -> None: + for stream in (sys.stdout, sys.stderr): + if hasattr(stream, "reconfigure"): + stream.reconfigure(encoding="utf-8", errors="replace") diff --git a/src/nlprog/command_safety.py b/src/nlprog/command_safety.py new file mode 100644 index 0000000..e921783 --- /dev/null +++ b/src/nlprog/command_safety.py @@ -0,0 +1,65 @@ +from __future__ import annotations + +import re +from dataclasses import dataclass + + +@dataclass(frozen=True) +class CommandAssessment: + level: str + reason: str + + @property + def blocked(self) -> bool: + return self.level == "block" + + @property + def needs_confirmation(self) -> bool: + return self.level == "confirm" + + +BLOCK_PATTERNS = [ + (r"\bformat\b", "Formatting disks is blocked."), + (r"\bdiskpart\b", "Disk partitioning is blocked."), + (r"\bshutdown\b", "System shutdown commands are blocked."), + (r"\brestart-computer\b", "System restart commands are blocked."), + (r"\bstop-computer\b", "System power commands are blocked."), + (r"\breg\s+(delete|add)\b", "Registry modification commands are blocked."), + (r"\bbcdedit\b", "Boot configuration commands are blocked."), + (r"\bnetsh\b", "Network stack modification commands are blocked."), + (r"\bcipher\s+/w\b", "Secure wipe commands are blocked."), +] + +CONFIRM_PATTERNS = [ + (r"\brm\b|\bdel\b|\berase\b|\bremove-item\b", "File deletion requires confirmation."), + (r"\brmdir\b|\brd\b", "Directory deletion requires confirmation."), + (r"\bmove-item\b|\bmv\b", "Moving files requires confirmation."), + (r"\bcopy-item\b|\bcp\b", "Copying files requires confirmation."), + (r"\bpip\s+install\b", "Installing Python packages requires confirmation."), + (r"\bnpm\s+install\b|\bnpm\s+i\b", "Installing npm packages requires confirmation."), + (r"\bpnpm\s+install\b|\byarn\s+add\b", "Installing Node packages requires confirmation."), + (r"\bcurl\b.*\|\s*(powershell|pwsh|sh|bash)", "Piped remote scripts require confirmation."), + (r"\birm\b.*\|\s*(iex|invoke-expression)", "Downloaded PowerShell execution requires confirmation."), + (r"\biwr\b.*\|\s*(iex|invoke-expression)", "Downloaded PowerShell execution requires confirmation."), + (r"\bset-executionpolicy\b", "Changing PowerShell execution policy requires confirmation."), +] + + +def assess_command(command: str) -> CommandAssessment: + text = _normalize(command) + if not text: + return CommandAssessment("block", "Empty command.") + + for pattern, reason in BLOCK_PATTERNS: + if re.search(pattern, text): + return CommandAssessment("block", reason) + + for pattern, reason in CONFIRM_PATTERNS: + if re.search(pattern, text): + return CommandAssessment("confirm", reason) + + return CommandAssessment("allow", "Command is allowed by the current safety policy.") + + +def _normalize(command: str) -> str: + return " ".join(command.lower().strip().split()) diff --git a/src/nlprog/config.py b/src/nlprog/config.py new file mode 100644 index 0000000..7bb52c5 --- /dev/null +++ b/src/nlprog/config.py @@ -0,0 +1,280 @@ +from __future__ import annotations + +import json +import os +from dataclasses import dataclass +from pathlib import Path +from typing import Any + + +CONFIG_DIR = Path.home() / ".nlprog" +CONFIG_FILE = CONFIG_DIR / "config.json" + +DEFAULT_CONFIG: dict[str, Any] = { + "provider": "mock", + "model": "mock-model", + "active_model": "mock", + "models": { + "mock": { + "provider": "mock", + "model": "mock-model", + } + }, + "api_key_env": "NLPROG_API_KEY", + "base_url": None, + "temperature": 0.2, + "require_confirmation": True, + "opencode_command": "opencode", + "json_repair_retries": 2, +} + +CONFIG_TYPES: dict[str, type] = { + "provider": str, + "model": str, + "api_key": str, + "api_key_env": str, + "base_url": str, + "temperature": float, + "require_confirmation": bool, + "max_context_files": int, + "timeout_seconds": int, + "opencode_command": str, + "json_repair_retries": int, +} + +MODEL_CONFIG_KEYS = { + "provider", + "model", + "api_key", + "api_key_env", + "base_url", + "temperature", + "timeout_seconds", + "opencode_command", + "json_repair_retries", +} + + +@dataclass(frozen=True) +class Config: + provider: str = "mock" + model: str = "mock-model" + api_key: str | None = None + api_key_env: str = "NLPROG_API_KEY" + base_url: str | None = None + temperature: float = 0.2 + require_confirmation: bool = True + max_context_files: int = 12 + timeout_seconds: int = 60 + opencode_command: str = "opencode" + json_repair_retries: int = 2 + + +def load_config() -> Config: + raw = _active_config() + + provider = os.getenv("NLPROG_PROVIDER", raw.get("provider", "mock")) + model = os.getenv("NLPROG_MODEL", raw.get("model", "mock-model")) + api_key_env = os.getenv("NLPROG_API_KEY_ENV", raw.get("api_key_env", "NLPROG_API_KEY")) + api_key = os.getenv(api_key_env) or os.getenv("NLPROG_API_KEY") or raw.get("api_key") + base_url = os.getenv("NLPROG_BASE_URL", raw.get("base_url")) + temperature = float(os.getenv("NLPROG_TEMPERATURE", raw.get("temperature", 0.2))) + require_confirmation = _bool_env( + os.getenv("NLPROG_REQUIRE_CONFIRMATION"), + bool(raw.get("require_confirmation", True)), + ) + opencode_command = os.getenv("NLPROG_OPENCODE_COMMAND", raw.get("opencode_command", "opencode")) + json_repair_retries = int(os.getenv("NLPROG_JSON_REPAIR_RETRIES", raw.get("json_repair_retries", 2))) + + return Config( + provider=provider, + model=model, + api_key=api_key, + api_key_env=api_key_env, + base_url=base_url, + temperature=temperature, + require_confirmation=require_confirmation, + max_context_files=int(raw.get("max_context_files", 12)), + timeout_seconds=int(raw.get("timeout_seconds", 60)), + opencode_command=opencode_command, + json_repair_retries=json_repair_retries, + ) + + +def write_default_config() -> Path: + CONFIG_DIR.mkdir(parents=True, exist_ok=True) + if not CONFIG_FILE.exists(): + write_config_file(dict(DEFAULT_CONFIG)) + return CONFIG_FILE + + +def read_config_file() -> dict[str, Any]: + if not CONFIG_FILE.exists(): + return {} + return json.loads(CONFIG_FILE.read_text(encoding="utf-8")) + + +def write_config_file(data: dict[str, Any]) -> Path: + CONFIG_DIR.mkdir(parents=True, exist_ok=True) + CONFIG_FILE.write_text(json.dumps(data, indent=2, ensure_ascii=False) + "\n", encoding="utf-8") + return CONFIG_FILE + + +def set_config_value(key: str, value: str) -> Path: + if key not in CONFIG_TYPES: + allowed = ", ".join(sorted(CONFIG_TYPES)) + raise ValueError(f"Unknown config key: {key}. Allowed keys: {allowed}") + data = dict(DEFAULT_CONFIG) + data.update(read_config_file()) + data[key] = parse_config_value(key, value) + return write_config_file(data) + + +def use_opencode_config(model: str = "", command: str = "opencode") -> Path: + data = dict(DEFAULT_CONFIG) + data.update(read_config_file()) + models = _models_from_data(data) + models["opencode"] = { + "provider": "opencode", + "model": model, + "opencode_command": command, + } + data["models"] = models + _apply_model(data, "opencode", models["opencode"]) + return write_config_file(data) + + +def list_models() -> tuple[str | None, dict[str, dict[str, Any]]]: + existing = read_config_file() + data = dict(DEFAULT_CONFIG) + data.update(existing) + active = data.get("active_model") + models = _models_from_data(existing) + if not models: + current = {key: data.get(key) for key in MODEL_CONFIG_KEYS if key in data} + return "current", {"current": current} + return active if isinstance(active, str) else None, models + + +def add_model( + name: str, + provider: str, + model: str = "", + base_url: str | None = None, + api_key_env: str | None = None, + opencode_command: str | None = None, + activate: bool = False, +) -> Path: + name = _normalize_model_name(name) + provider = provider.strip() + if not provider: + raise ValueError("Model provider cannot be empty.") + + data = dict(DEFAULT_CONFIG) + data.update(read_config_file()) + models = _models_from_data(data) + entry: dict[str, Any] = { + "provider": provider, + "model": model, + } + if base_url: + entry["base_url"] = base_url + if api_key_env: + entry["api_key_env"] = api_key_env + if opencode_command: + entry["opencode_command"] = opencode_command + + models[name] = entry + data["models"] = models + if activate: + _apply_model(data, name, entry) + return write_config_file(data) + + +def use_model(name: str) -> Path: + name = _normalize_model_name(name) + data = dict(DEFAULT_CONFIG) + data.update(read_config_file()) + models = _models_from_data(data) + if name not in models: + raise ValueError(f"Unknown model: {name}") + _apply_model(data, name, models[name]) + data["models"] = models + return write_config_file(data) + + +def remove_model(name: str) -> Path: + name = _normalize_model_name(name) + data = dict(DEFAULT_CONFIG) + data.update(read_config_file()) + models = _models_from_data(data) + if name not in models: + raise ValueError(f"Unknown model: {name}") + if data.get("active_model") == name: + raise ValueError("Cannot remove the active model. Use another model first.") + del models[name] + data["models"] = models + return write_config_file(data) + + +def parse_config_value(key: str, value: str) -> Any: + expected = CONFIG_TYPES[key] + stripped = value.strip() + lowered = stripped.lower() + if lowered in {"null", "none"}: + return None + if expected is bool: + if lowered in {"1", "true", "yes", "on"}: + return True + if lowered in {"0", "false", "no", "off"}: + return False + raise ValueError(f"Expected boolean value for {key}.") + if expected is int: + return int(stripped) + if expected is float: + return float(stripped) + return value + + +def _bool_env(value: str | None, default: bool) -> bool: + if value is None: + return default + return value.strip().lower() in {"1", "true", "yes", "on"} + + +def _active_config() -> dict[str, Any]: + existing = read_config_file() + raw = dict(DEFAULT_CONFIG) + raw.update(existing) + active = raw.get("active_model") + models = _models_from_data(raw) + if "active_model" in existing and isinstance(active, str) and isinstance(models.get(active), dict): + for key, value in models[active].items(): + if key in MODEL_CONFIG_KEYS: + raw[key] = value + return raw + + +def _models_from_data(data: dict[str, Any]) -> dict[str, dict[str, Any]]: + models = data.get("models") + if not isinstance(models, dict): + return {} + normalized: dict[str, dict[str, Any]] = {} + for name, value in models.items(): + if isinstance(name, str) and isinstance(value, dict): + normalized[name] = dict(value) + return normalized + + +def _apply_model(data: dict[str, Any], name: str, entry: dict[str, Any]) -> None: + data["active_model"] = name + for key, value in entry.items(): + if key in MODEL_CONFIG_KEYS: + data[key] = value + + +def _normalize_model_name(name: str) -> str: + normalized = name.strip() + if not normalized: + raise ValueError("Model name cannot be empty.") + return normalized diff --git a/src/nlprog/context.py b/src/nlprog/context.py new file mode 100644 index 0000000..618b2b1 --- /dev/null +++ b/src/nlprog/context.py @@ -0,0 +1,69 @@ +from __future__ import annotations + +from pathlib import Path + + +SKIP_DIRS = { + ".git", + ".venv", + "__pycache__", + "node_modules", + "dist", + "build", + ".mypy_cache", + ".pytest_cache", +} + +TEXT_EXTENSIONS = { + ".py", + ".js", + ".ts", + ".tsx", + ".jsx", + ".json", + ".toml", + ".yaml", + ".yml", + ".md", + ".txt", + ".css", + ".html", + ".java", + ".go", + ".rs", + ".cs", + ".php", + ".rb", + ".sh", + ".ps1", +} + + +def collect_context(root: Path, max_files: int) -> str: + files: list[Path] = [] + for path in root.rglob("*"): + if len(files) >= max_files: + break + if any(part in SKIP_DIRS for part in path.parts): + continue + if path.is_file() and path.suffix.lower() in TEXT_EXTENSIONS: + files.append(path) + + if not files: + return "No project files were found." + + chunks: list[str] = [] + for path in files: + rel = path.relative_to(root) + text = _safe_read(path) + chunks.append(f"### {rel}\n{text[:4000]}") + return "\n\n".join(chunks) + + +def _safe_read(path: Path) -> str: + try: + return path.read_text(encoding="utf-8") + except UnicodeDecodeError: + return path.read_text(encoding="utf-8", errors="replace") + except OSError as exc: + return f"[Could not read file: {exc}]" diff --git a/src/nlprog/doctor.py b/src/nlprog/doctor.py new file mode 100644 index 0000000..e400b5c --- /dev/null +++ b/src/nlprog/doctor.py @@ -0,0 +1,117 @@ +from __future__ import annotations + +import shutil +import subprocess +import sys +from dataclasses import dataclass +from pathlib import Path + +from .config import CONFIG_FILE, load_config +from .encoding_utils import decode_process_output +from .project_profile import load_project_profile + + +@dataclass(frozen=True) +class DoctorCheck: + status: str + name: str + detail: str + + +def run_doctor(root: Path, check_opencode: bool = True) -> list[DoctorCheck]: + root = root.resolve() + config = load_config() + profile = load_project_profile(root) + checks: list[DoctorCheck] = [] + + checks.append(DoctorCheck("OK", "Python", sys.version.split()[0])) + checks.append(_path_check("Package", root / "src" / "nlprog" / "__init__.py", "Source package found.")) + checks.append( + DoctorCheck( + "OK" if CONFIG_FILE.exists() else "WARN", + "User config", + str(CONFIG_FILE) if CONFIG_FILE.exists() else "Run `python -m nlprog init` to create one.", + ) + ) + checks.append(DoctorCheck("OK", "Provider", f"{config.provider} / {config.model or '(default)'}")) + checks.append( + DoctorCheck( + "OK" if config.json_repair_retries >= 0 else "FAIL", + "JSON repair", + f"Retries: {config.json_repair_retries}", + ) + ) + + project_file = root / ".nlprog" / "project.json" + checks.append(_path_check("Project profile", project_file, "Project profile exists.")) + checks.append( + DoctorCheck( + "OK" if profile.verification_commands else "WARN", + "Verification", + ", ".join(profile.verification_commands) if profile.verification_commands else "No verification command detected.", + ) + ) + checks.append( + DoctorCheck( + "OK" if profile.protected_paths else "WARN", + "Protected paths", + ", ".join(profile.protected_paths) if profile.protected_paths else "No protected paths configured.", + ) + ) + + if check_opencode: + checks.extend(_check_opencode(config.opencode_command)) + + return checks + + +def _path_check(name: str, path: Path, ok_detail: str) -> DoctorCheck: + if path.exists(): + return DoctorCheck("OK", name, ok_detail) + return DoctorCheck("WARN", name, f"Missing: {path}") + + +def _check_opencode(command: str) -> list[DoctorCheck]: + checks: list[DoctorCheck] = [] + found = shutil.which(command) + checks.append( + DoctorCheck( + "OK" if found else "WARN", + "opencode command", + found or f"Not found in PATH. Set NLPROG_OPENCODE_COMMAND if needed.", + ) + ) + auth_path = Path.home() / ".local" / "share" / "opencode" / "auth.json" + checks.append( + DoctorCheck( + "OK" if auth_path.exists() else "WARN", + "opencode auth", + str(auth_path) if auth_path.exists() else "Auth file not found. Run opencode /connect.", + ) + ) + if not found: + return checks + + try: + completed = subprocess.run( + [command, "--version"], + capture_output=True, + timeout=15, + ) + except Exception as exc: + checks.append(DoctorCheck("WARN", "opencode version", f"Could not run opencode: {exc}")) + return checks + + if completed.returncode == 0: + detail = ( + decode_process_output(completed.stdout, prefer_utf8=True) + or decode_process_output(completed.stderr, prefer_utf8=True) + ).strip() + checks.append(DoctorCheck("OK", "opencode version", detail)) + else: + detail = ( + decode_process_output(completed.stderr, prefer_utf8=True) + or decode_process_output(completed.stdout, prefer_utf8=True) + ).strip() + checks.append(DoctorCheck("WARN", "opencode version", detail or f"Exit code {completed.returncode}")) + return checks diff --git a/src/nlprog/encoding_utils.py b/src/nlprog/encoding_utils.py new file mode 100644 index 0000000..802276c --- /dev/null +++ b/src/nlprog/encoding_utils.py @@ -0,0 +1,27 @@ +from __future__ import annotations + +import locale + + +def decode_process_output(data: bytes | None, prefer_utf8: bool = True) -> str: + if not data: + return "" + + encodings: list[str] = [] + if prefer_utf8: + encodings.append("utf-8") + preferred = locale.getpreferredencoding(False) + encodings.extend([preferred, "utf-8", "gbk", "mbcs"]) + + seen: set[str] = set() + for encoding in encodings: + normalized = encoding.lower() + if normalized in seen: + continue + seen.add(normalized) + try: + return data.decode(encoding) + except (LookupError, UnicodeDecodeError): + continue + + return data.decode(preferred, errors="replace") diff --git a/src/nlprog/executor.py b/src/nlprog/executor.py new file mode 100644 index 0000000..cc7066c --- /dev/null +++ b/src/nlprog/executor.py @@ -0,0 +1,41 @@ +from __future__ import annotations + +import os +import subprocess +from dataclasses import dataclass +from pathlib import Path + +from .agent import CommandStep +from .encoding_utils import decode_process_output + + +@dataclass(frozen=True) +class CommandResult: + cmd: str + returncode: int + stdout: str + stderr: str + + +def run_command(step: CommandStep, cwd: Path, timeout: int) -> CommandResult: + completed = subprocess.run( + step.cmd, + cwd=str(cwd), + shell=True, + capture_output=True, + env=_subprocess_env(), + timeout=timeout, + ) + return CommandResult( + cmd=step.cmd, + returncode=completed.returncode, + stdout=decode_process_output(completed.stdout), + stderr=decode_process_output(completed.stderr), + ) + + +def _subprocess_env() -> dict[str, str]: + env = os.environ.copy() + env.setdefault("PYTHONIOENCODING", "utf-8") + env.setdefault("PYTHONUTF8", "1") + return env diff --git a/src/nlprog/json_repair.py b/src/nlprog/json_repair.py new file mode 100644 index 0000000..4fdecc4 --- /dev/null +++ b/src/nlprog/json_repair.py @@ -0,0 +1,83 @@ +from __future__ import annotations + +import json +from dataclasses import dataclass +from typing import Any, Callable + +from .llm import LLMClient, Message + + +@dataclass(frozen=True) +class JsonResult: + data: dict[str, Any] + raw: str + attempts: int + repaired: bool = False + + +def complete_json( + client: LLMClient, + messages: list[Message], + schema_hint: str, + validate: Callable[[dict[str, Any]], str | None], + max_retries: int, +) -> JsonResult: + working_messages = list(messages) + raw = "" + last_error = "" + + for attempt in range(max_retries + 1): + raw = client.complete(working_messages) + parsed, error = parse_json_object(raw) + if parsed is not None: + validation_error = validate(parsed) + if validation_error is None: + return JsonResult(parsed, raw, attempt + 1, repaired=attempt > 0) + last_error = validation_error + else: + last_error = error + + working_messages.append(Message("assistant", raw)) + working_messages.append( + Message( + "user", + "\n".join( + [ + "Your previous response was not valid for this interface.", + f"Problem: {last_error}", + "Return only a single JSON object. Do not include markdown, prose, comments, or code fences.", + "Required schema:", + schema_hint, + ] + ), + ) + ) + + return JsonResult( + {"final": f"The model did not return valid JSON after {max_retries + 1} attempts. Last problem: {last_error}"}, + raw, + max_retries + 1, + repaired=max_retries > 0, + ) + + +def parse_json_object(raw: str) -> tuple[dict[str, Any] | None, str]: + text = strip_fences(raw.strip()) + try: + data = json.loads(text) + except json.JSONDecodeError as exc: + return None, f"JSON parse error at line {exc.lineno}, column {exc.colno}: {exc.msg}" + if not isinstance(data, dict): + return None, "Top-level JSON value must be an object." + return data, "" + + +def strip_fences(text: str) -> str: + if text.startswith("```"): + lines = text.splitlines() + if lines and lines[0].startswith("```"): + lines = lines[1:] + if lines and lines[-1].startswith("```"): + lines = lines[:-1] + return "\n".join(lines) + return text diff --git a/src/nlprog/llm.py b/src/nlprog/llm.py new file mode 100644 index 0000000..49d9c54 --- /dev/null +++ b/src/nlprog/llm.py @@ -0,0 +1,234 @@ +from __future__ import annotations + +import json +import re +import subprocess +import urllib.error +import urllib.request +from dataclasses import dataclass +from typing import Any, Protocol + +from .config import Config +from .encoding_utils import decode_process_output + + +class LLMError(RuntimeError): + pass + + +@dataclass(frozen=True) +class Message: + role: str + content: str + + +class LLMClient(Protocol): + def complete(self, messages: list[Message]) -> str: + raise NotImplementedError + + +def create_client(config: Config) -> LLMClient: + provider = config.provider.lower() + if provider == "mock": + return MockClient() + if provider in {"openai", "openai-compatible"}: + base_url = config.base_url or "https://api.openai.com/v1" + return OpenAICompatibleClient(config, base_url) + if provider == "anthropic": + return AnthropicClient(config) + if provider == "gemini": + return GeminiClient(config) + if provider == "opencode": + return OpencodeClient(config) + raise LLMError(f"Unsupported provider: {config.provider}") + + +class MockClient: + def complete(self, messages: list[Message]) -> str: + if messages and ( + "long-term project memory" in messages[-1].content + or "stable project memory" in messages[-1].content + ): + return json.dumps( + {"items": ["Use `python -m compileall src` as the focused syntax verification command."]}, + ensure_ascii=False, + ) + + if messages and "NLProg Agent" in messages[0].content: + if any("Tool result" in message.content for message in messages): + return json.dumps( + {"final": "Mock agent inspected the project and finished without making changes."}, + ensure_ascii=False, + ) + if "patch" in messages[-1].content.lower() or "补丁" in messages[-1].content: + return json.dumps( + { + "action": "apply_patch", + "args": { + "patch": "*** Begin Patch\n*** Add File: tmp/mock-patch.txt\n+hello from patch\n*** End Patch" + }, + "reason": "Exercise the patch preview and apply flow.", + }, + ensure_ascii=False, + ) + if any("Tool result" in message.content for message in messages): + return json.dumps( + {"final": "Mock agent inspected the project and finished without making changes."}, + ensure_ascii=False, + ) + if "create file" in messages[-1].content.lower() or "创建文件" in messages[-1].content: + return json.dumps( + { + "action": "create_file", + "args": {"path": "tmp/mock-agent-preview.txt", "content": "hello from mock agent\n"}, + "reason": "Create a small file to exercise edit preview and confirmation.", + }, + ensure_ascii=False, + ) + return json.dumps( + { + "action": "list_files", + "args": {"path": "."}, + "reason": "Inspect the project root before deciding what to do.", + }, + ensure_ascii=False, + ) + + user_text = messages[-1].content + command = "dir" if "列出" in user_text or "list" in user_text.lower() else "echo Mock mode: no real model was called" + return json.dumps( + { + "summary": "Mock provider generated a safe demonstration command.", + "commands": [{"cmd": command, "reason": "Demonstrate the execution flow."}], + "notes": ["Set NLPROG_PROVIDER and NLPROG_API_KEY to use a real model."], + }, + ensure_ascii=False, + ) + + +class OpenAICompatibleClient: + def __init__(self, config: Config, base_url: str) -> None: + self.config = config + self.base_url = base_url.rstrip("/") + + def complete(self, messages: list[Message]) -> str: + if not self.config.api_key: + raise LLMError("Missing API key. Set NLPROG_API_KEY or api_key in ~/.nlprog/config.json.") + + payload = { + "model": self.config.model, + "messages": [{"role": m.role, "content": m.content} for m in messages], + "temperature": self.config.temperature, + } + data = _post_json( + f"{self.base_url}/chat/completions", + payload, + { + "Authorization": f"Bearer {self.config.api_key}", + "Content-Type": "application/json", + }, + self.config.timeout_seconds, + ) + return data["choices"][0]["message"]["content"] + + +class AnthropicClient: + def __init__(self, config: Config) -> None: + self.config = config + + def complete(self, messages: list[Message]) -> str: + if not self.config.api_key: + raise LLMError("Missing API key. Set NLPROG_API_KEY or api_key in ~/.nlprog/config.json.") + + system = "\n\n".join(m.content for m in messages if m.role == "system") + user_messages = [{"role": m.role, "content": m.content} for m in messages if m.role != "system"] + payload = { + "model": self.config.model, + "system": system, + "messages": user_messages, + "max_tokens": 2048, + "temperature": self.config.temperature, + } + data = _post_json( + "https://api.anthropic.com/v1/messages", + payload, + { + "x-api-key": self.config.api_key, + "anthropic-version": "2023-06-01", + "Content-Type": "application/json", + }, + self.config.timeout_seconds, + ) + parts = data.get("content", []) + return "\n".join(part.get("text", "") for part in parts if part.get("type") == "text") + + +class GeminiClient: + def __init__(self, config: Config) -> None: + self.config = config + + def complete(self, messages: list[Message]) -> str: + if not self.config.api_key: + raise LLMError("Missing API key. Set NLPROG_API_KEY or api_key in ~/.nlprog/config.json.") + + prompt = "\n\n".join(f"{m.role.upper()}:\n{m.content}" for m in messages) + payload = { + "contents": [{"parts": [{"text": prompt}]}], + "generationConfig": {"temperature": self.config.temperature}, + } + url = f"https://generativelanguage.googleapis.com/v1beta/models/{self.config.model}:generateContent?key={self.config.api_key}" + data = _post_json(url, payload, {"Content-Type": "application/json"}, self.config.timeout_seconds) + return data["candidates"][0]["content"]["parts"][0]["text"] + + +class OpencodeClient: + def __init__(self, config: Config) -> None: + self.config = config + + def complete(self, messages: list[Message]) -> str: + prompt = "\n\n".join(f"{m.role.upper()}:\n{m.content}" for m in messages) + command = [self.config.opencode_command, "run"] + if self.config.model and self.config.model != "mock-model": + command.extend(["--model", self.config.model]) + command.append(prompt) + + try: + completed = subprocess.run( + command, + capture_output=True, + timeout=self.config.timeout_seconds, + ) + except FileNotFoundError as exc: + raise LLMError( + f"opencode was not found. Install it or set NLPROG_OPENCODE_COMMAND. Tried: {self.config.opencode_command}" + ) from exc + except subprocess.TimeoutExpired as exc: + raise LLMError(f"opencode timed out after {self.config.timeout_seconds} seconds.") from exc + + if completed.returncode != 0: + stderr = _strip_ansi(decode_process_output(completed.stderr, prefer_utf8=True).strip()) + stdout = _strip_ansi(decode_process_output(completed.stdout, prefer_utf8=True).strip()) + detail = stderr or stdout or f"exit code {completed.returncode}" + raise LLMError(f"opencode failed: {detail}") + return _strip_ansi(decode_process_output(completed.stdout, prefer_utf8=True).strip()) + + +def _post_json(url: str, payload: dict[str, Any], headers: dict[str, str], timeout: int) -> dict[str, Any]: + request = urllib.request.Request( + url, + data=json.dumps(payload).encode("utf-8"), + headers=headers, + method="POST", + ) + try: + with urllib.request.urlopen(request, timeout=timeout) as response: + return json.loads(response.read().decode("utf-8")) + except urllib.error.HTTPError as exc: + body = exc.read().decode("utf-8", errors="replace") + raise LLMError(f"LLM request failed with HTTP {exc.code}: {body}") from exc + except urllib.error.URLError as exc: + raise LLMError(f"LLM request failed: {exc}") from exc + + +def _strip_ansi(text: str) -> str: + return re.sub(r"\x1b\[[0-9;?]*[ -/]*[@-~]", "", text) diff --git a/src/nlprog/patching.py b/src/nlprog/patching.py new file mode 100644 index 0000000..7684fd8 --- /dev/null +++ b/src/nlprog/patching.py @@ -0,0 +1,158 @@ +from __future__ import annotations + +from dataclasses import dataclass +from pathlib import Path + + +class PatchError(ValueError): + pass + + +@dataclass(frozen=True) +class PatchOp: + kind: str + path: str + old: str = "" + new: str = "" + + +def parse_patch(patch: str) -> list[PatchOp]: + lines = patch.splitlines() + if not lines or lines[0].strip() != "*** Begin Patch": + raise PatchError("Patch must start with *** Begin Patch.") + if lines[-1].strip() != "*** End Patch": + raise PatchError("Patch must end with *** End Patch.") + + ops: list[PatchOp] = [] + index = 1 + while index < len(lines) - 1: + line = lines[index] + if line.startswith("*** Add File: "): + path = _remove_prefix(line, "*** Add File: ").strip() + index += 1 + new_lines: list[str] = [] + while index < len(lines) - 1 and not lines[index].startswith("*** "): + if not lines[index].startswith("+"): + raise PatchError(f"Add File lines must start with '+': {path}") + new_lines.append(lines[index][1:]) + index += 1 + ops.append(PatchOp("add", path, new="\n".join(new_lines) + ("\n" if new_lines else ""))) + continue + + if line.startswith("*** Delete File: "): + path = _remove_prefix(line, "*** Delete File: ").strip() + ops.append(PatchOp("delete", path)) + index += 1 + continue + + if line.startswith("*** Update File: "): + path = _remove_prefix(line, "*** Update File: ").strip() + index += 1 + old_lines: list[str] = [] + new_lines: list[str] = [] + while index < len(lines) - 1 and not lines[index].startswith("*** "): + current = lines[index] + if current.startswith("@@"): + index += 1 + continue + if not current: + old_lines.append("") + new_lines.append("") + elif current[0] == " ": + old_lines.append(current[1:]) + new_lines.append(current[1:]) + elif current[0] == "-": + old_lines.append(current[1:]) + elif current[0] == "+": + new_lines.append(current[1:]) + else: + raise PatchError(f"Update lines must start with space, '+', '-', or '@@': {path}") + index += 1 + ops.append(PatchOp("update", path, old=_join_patch_lines(old_lines), new=_join_patch_lines(new_lines))) + continue + + raise PatchError(f"Unknown patch header: {line}") + + if not ops: + raise PatchError("Patch contains no operations.") + return ops + + +def apply_patch_to_root(root: Path, patch: str, protected_paths: set[str]) -> list[str]: + root = root.resolve() + ops = parse_patch(patch) + messages: list[str] = [] + for op in ops: + target = _resolve(root, op.path) + if _is_protected(root, target, protected_paths): + raise PatchError(f"Refusing to modify protected path: {op.path}") + + if op.kind == "add": + if target.exists(): + raise PatchError(f"Cannot add existing file: {op.path}") + target.parent.mkdir(parents=True, exist_ok=True) + target.write_text(op.new, encoding="utf-8") + messages.append(f"Added {op.path}") + elif op.kind == "delete": + if not target.is_file(): + raise PatchError(f"Cannot delete missing file: {op.path}") + target.unlink() + messages.append(f"Deleted {op.path}") + elif op.kind == "update": + if not target.is_file(): + raise PatchError(f"Cannot update missing file: {op.path}") + text = target.read_text(encoding="utf-8", errors="replace") + count = text.count(op.old) + if count != 1: + raise PatchError(f"Expected exactly one match in {op.path}, found {count}.") + target.write_text(text.replace(op.old, op.new, 1), encoding="utf-8") + messages.append(f"Updated {op.path}") + else: + raise PatchError(f"Unsupported operation: {op.kind}") + return messages + + +def summarize_patch(patch: str) -> str: + ops = parse_patch(patch) + rows: list[str] = [] + for op in ops: + if op.kind == "add": + rows.append(f"Add {op.path} ({len(op.new)} bytes)") + elif op.kind == "delete": + rows.append(f"Delete {op.path}") + elif op.kind == "update": + rows.append(f"Update {op.path} (-{_line_count(op.old)} +{_line_count(op.new)} lines)") + return "\n".join(rows) + + +def _join_patch_lines(lines: list[str]) -> str: + return "\n".join(lines) + ("\n" if lines else "") + + +def _remove_prefix(text: str, prefix: str) -> str: + if text.startswith(prefix): + return text[len(prefix) :] + return text + + +def _line_count(text: str) -> int: + if not text: + return 0 + return len(text.splitlines()) + + +def _resolve(root: Path, path: str) -> Path: + target = (root / path).resolve() + if target != root and root not in target.parents: + raise PatchError(f"Path escapes workspace: {path}") + return target + + +def _is_protected(root: Path, target: Path, protected_paths: set[str]) -> bool: + try: + rel = target.relative_to(root) + except ValueError: + return True + parts = rel.parts + rel_text = str(rel) + return any(protected in parts or rel_text.startswith(protected.rstrip("/\\") + "\\") for protected in protected_paths) diff --git a/src/nlprog/project_profile.py b/src/nlprog/project_profile.py new file mode 100644 index 0000000..6c4519b --- /dev/null +++ b/src/nlprog/project_profile.py @@ -0,0 +1,223 @@ +from __future__ import annotations + +import json +from datetime import datetime +from dataclasses import dataclass, field +from pathlib import Path +from typing import Any + + +PROJECT_DIR = ".nlprog" +PROJECT_FILE = "project.json" +RULES_FILE = "rules.md" +MEMORY_FILE = "memory.md" + + +@dataclass(frozen=True) +class ProjectProfile: + project_types: list[str] = field(default_factory=list) + important_files: list[str] = field(default_factory=list) + verification_commands: list[str] = field(default_factory=list) + protected_paths: list[str] = field(default_factory=list) + rules: str = "" + memory: str = "" + source: str = "inferred" + + +def init_project(root: Path, force: bool = False) -> list[Path]: + root = root.resolve() + profile = infer_project(root) + nlprog_dir = root / PROJECT_DIR + nlprog_dir.mkdir(exist_ok=True) + + created: list[Path] = [] + project_path = nlprog_dir / PROJECT_FILE + rules_path = nlprog_dir / RULES_FILE + memory_path = nlprog_dir / MEMORY_FILE + + project_data = { + "schema_version": 1, + "project_types": profile.project_types, + "important_files": profile.important_files, + "verification_commands": profile.verification_commands, + "protected_paths": profile.protected_paths, + } + _write_if_needed(project_path, json.dumps(project_data, indent=2, ensure_ascii=False) + "\n", force, created) + + rules = _default_rules(profile) + _write_if_needed(rules_path, rules, force, created) + + memory = "\n".join( + [ + "# NLProg Memory", + "", + "- Add stable project facts here as the agent learns them.", + "- Keep secrets, tokens, and private credentials out of this file.", + "", + ] + ) + _write_if_needed(memory_path, memory, force, created) + return created + + +def load_project_profile(root: Path) -> ProjectProfile: + root = root.resolve() + inferred = infer_project(root) + nlprog_dir = root / PROJECT_DIR + project_path = nlprog_dir / PROJECT_FILE + rules_path = nlprog_dir / RULES_FILE + memory_path = nlprog_dir / MEMORY_FILE + + data: dict[str, Any] = {} + source = "inferred" + if project_path.exists(): + try: + data = json.loads(project_path.read_text(encoding="utf-8")) + source = str(project_path.relative_to(root)) + except (OSError, json.JSONDecodeError): + data = {} + + rules = _read_optional(rules_path) + memory = _read_optional(memory_path) + + return ProjectProfile( + project_types=_list_or_default(data.get("project_types"), inferred.project_types), + important_files=_list_or_default(data.get("important_files"), inferred.important_files), + verification_commands=_list_or_default(data.get("verification_commands"), inferred.verification_commands), + protected_paths=_list_or_default(data.get("protected_paths"), inferred.protected_paths), + rules=rules, + memory=memory, + source=source, + ) + + +def build_project_context(root: Path) -> str: + profile = load_project_profile(root) + lines = [ + f"Project profile source: {profile.source}", + f"Project types: {', '.join(profile.project_types) or 'unknown'}", + "Important files:", + _format_list(profile.important_files), + "Suggested verification commands:", + _format_list(profile.verification_commands), + "Protected paths:", + _format_list(profile.protected_paths), + ] + if profile.rules.strip(): + lines.extend(["Project rules:", profile.rules.strip()]) + if profile.memory.strip(): + lines.extend(["Project memory:", profile.memory.strip()]) + if profile.source == "inferred": + lines.append("Tip: run `python -m nlprog init-project` to save editable project rules and memory.") + return "\n".join(lines) + + +def append_memory(root: Path, items: list[str]) -> Path: + root = root.resolve() + nlprog_dir = root / PROJECT_DIR + nlprog_dir.mkdir(exist_ok=True) + memory_path = nlprog_dir / MEMORY_FILE + if not memory_path.exists(): + memory_path.write_text("# NLProg Memory\n\n", encoding="utf-8") + + date = datetime.now().strftime("%Y-%m-%d") + existing = memory_path.read_text(encoding="utf-8", errors="replace") + lines = [existing.rstrip(), "", f"## {date}", ""] + lines.extend(f"- {item}" for item in items) + lines.append("") + memory_path.write_text("\n".join(lines), encoding="utf-8") + return memory_path + + +def infer_project(root: Path) -> ProjectProfile: + project_types: list[str] = [] + important_files: list[str] = [] + verification_commands: list[str] = [] + + markers = { + "pyproject.toml": ("python", "python -m compileall src"), + "requirements.txt": ("python", "python -m compileall ."), + "package.json": ("node", "npm test"), + "pnpm-lock.yaml": ("node", "pnpm test"), + "yarn.lock": ("node", "yarn test"), + "go.mod": ("go", "go test ./..."), + "Cargo.toml": ("rust", "cargo test"), + } + + for filename, (project_type, verify_command) in markers.items(): + if (root / filename).exists(): + important_files.append(filename) + if project_type not in project_types: + project_types.append(project_type) + if verify_command not in verification_commands: + verification_commands.append(verify_command) + + if (root / "pytest.ini").exists() or (root / "tests").exists(): + if "python" not in project_types: + project_types.append("python") + if "python -m pytest" not in verification_commands: + verification_commands.append("python -m pytest") + + if (root / "src").exists() and "python" in project_types and "python -m compileall src" not in verification_commands: + verification_commands.insert(0, "python -m compileall src") + + return ProjectProfile( + project_types=project_types, + important_files=important_files, + verification_commands=verification_commands, + protected_paths=[".git", ".venv", "node_modules", "__pycache__", "dist", "build"], + ) + + +def _default_rules(profile: ProjectProfile) -> str: + commands = "\n".join(f"- `{command}`" for command in profile.verification_commands) or "- Add one here." + protected = "\n".join(f"- `{path}`" for path in profile.protected_paths) or "- Add one here." + return "\n".join( + [ + "# NLProg Rules", + "", + "## Editing", + "", + "- Inspect relevant files before editing.", + "- Prefer small, exact replacements over whole-file rewrites.", + "- Preview edits and ask for confirmation before writing files.", + "- Do not write secrets, API keys, or private credentials into the repository.", + "", + "## Protected Paths", + "", + protected, + "", + "## Verification", + "", + commands, + "", + ] + ) + + +def _write_if_needed(path: Path, content: str, force: bool, created: list[Path]) -> None: + if path.exists() and not force: + return + path.write_text(content, encoding="utf-8") + created.append(path) + + +def _read_optional(path: Path) -> str: + if not path.exists(): + return "" + try: + return path.read_text(encoding="utf-8") + except OSError: + return "" + + +def _list_or_default(value: object, default: list[str]) -> list[str]: + if isinstance(value, list): + return [str(item) for item in value] + return default + + +def _format_list(items: list[str]) -> str: + if not items: + return "- none" + return "\n".join(f"- {item}" for item in items) diff --git a/src/nlprog/run_log.py b/src/nlprog/run_log.py new file mode 100644 index 0000000..f4d5761 --- /dev/null +++ b/src/nlprog/run_log.py @@ -0,0 +1,58 @@ +from __future__ import annotations + +import json +from datetime import datetime +from pathlib import Path +from typing import Any + + +def write_run_log(root: Path, task: str, final: str, events: list[object], started_at: str) -> Path: + root = root.resolve() + runs_dir = root / ".nlprog" / "runs" + runs_dir.mkdir(parents=True, exist_ok=True) + ended_at = _now() + filename = ended_at.replace(":", "").replace("-", "").replace("T", "-") + ".json" + path = runs_dir / filename + payload: dict[str, Any] = { + "task": task, + "started_at": started_at, + "ended_at": ended_at, + "final": final, + "events": [{"kind": getattr(event, "kind", ""), "message": getattr(event, "message", "")} for event in events], + } + path.write_text(json.dumps(payload, indent=2, ensure_ascii=False) + "\n", encoding="utf-8") + return path + + +def list_run_logs(root: Path) -> list[Path]: + runs_dir = root.resolve() / ".nlprog" / "runs" + if not runs_dir.exists(): + return [] + return sorted(runs_dir.glob("*.json"), key=lambda path: path.name, reverse=True) + + +def load_run_log(root: Path, run_id: str) -> tuple[Path, dict[str, Any]]: + logs = list_run_logs(root) + if not logs: + raise FileNotFoundError("No run logs found.") + + if run_id == "latest": + path = logs[0] + else: + matches = [path for path in logs if path.stem == run_id or path.name == run_id or path.stem.startswith(run_id)] + if not matches: + raise FileNotFoundError(f"No run log matches: {run_id}") + if len(matches) > 1: + choices = ", ".join(path.stem for path in matches[:5]) + raise ValueError(f"Run id is ambiguous. Matches: {choices}") + path = matches[0] + + return path, json.loads(path.read_text(encoding="utf-8")) + + +def now_timestamp() -> str: + return _now() + + +def _now() -> str: + return datetime.now().replace(microsecond=0).isoformat() diff --git a/src/nlprog/tools.py b/src/nlprog/tools.py new file mode 100644 index 0000000..fe12f19 --- /dev/null +++ b/src/nlprog/tools.py @@ -0,0 +1,260 @@ +from __future__ import annotations + +import os +import subprocess +from dataclasses import dataclass +from pathlib import Path + +from .command_safety import assess_command +from .encoding_utils import decode_process_output +from .patching import PatchError, apply_patch_to_root, summarize_patch + + +SKIP_DIRS = {".git", ".venv", "__pycache__", "node_modules", "dist", "build"} + + +@dataclass(frozen=True) +class ToolResult: + ok: bool + output: str + + +@dataclass(frozen=True) +class PendingEdit: + action: str + path: str + preview: str + args: dict[str, object] + + +class ToolBox: + def __init__( + self, + root: Path, + timeout_seconds: int, + protected_paths: list[str] | None = None, + auto_approve_commands: bool = False, + ) -> None: + self.root = root.resolve() + self.timeout_seconds = timeout_seconds + self.auto_approve_commands = auto_approve_commands + self.protected_paths = set(SKIP_DIRS) + if protected_paths: + self.protected_paths.update(protected_paths) + + def run(self, name: str, args: dict[str, object], approved: bool = False) -> ToolResult: + if name == "list_files": + return self.list_files(str(args.get("path", "."))) + if name == "read_file": + return self.read_file(str(args.get("path", ""))) + if name == "search_text": + return self.search_text(str(args.get("query", "")), str(args.get("path", "."))) + if name == "create_file": + return self.create_file(str(args.get("path", "")), str(args.get("content", ""))) + if name == "replace_in_file": + return self.replace_in_file( + str(args.get("path", "")), + str(args.get("old", "")), + str(args.get("new", "")), + ) + if name == "run_command": + return self.run_command(str(args.get("cmd", "")), approved=approved) + if name == "apply_patch": + return self.apply_patch(str(args.get("patch", ""))) + return ToolResult(False, f"Unknown tool: {name}") + + def preview_edit(self, name: str, args: dict[str, object]) -> PendingEdit | None: + if name == "create_file": + path = str(args.get("path", "")) + content = str(args.get("content", "")) + target = self._resolve(path) + preview = "\n".join( + [ + f"Create file: {target.relative_to(self.root)}", + f"Bytes: {len(content)}", + "", + _preview_text(content), + ] + ) + return PendingEdit(name, path, preview, args) + + if name == "replace_in_file": + path = str(args.get("path", "")) + old = str(args.get("old", "")) + new = str(args.get("new", "")) + target = self._resolve(path) + if not target.is_file(): + return PendingEdit(name, path, f"Cannot preview: not a file: {path}", args) + text = target.read_text(encoding="utf-8", errors="replace") + count = text.count(old) if old else 0 + preview = "\n".join( + [ + f"Update file: {target.relative_to(self.root)}", + f"Matches: {count}", + "", + "--- old", + _preview_text(old), + "--- new", + _preview_text(new), + ] + ) + return PendingEdit(name, path, preview, args) + + if name == "apply_patch": + patch = str(args.get("patch", "")) + try: + summary = summarize_patch(patch) + except PatchError as exc: + summary = f"Patch parse error: {exc}" + preview = "\n".join(["Patch operations:", summary, "", _preview_text(patch, limit=12000)]) + return PendingEdit(name, "", preview, args) + + if name == "run_command": + cmd = str(args.get("cmd", "")) + assessment = assess_command(cmd) + if assessment.needs_confirmation: + preview = "\n".join(["Command requires confirmation:", assessment.reason, "", cmd]) + return PendingEdit(name, "", preview, args) + + return None + + def list_files(self, path: str) -> ToolResult: + target = self._resolve(path) + if not target.exists(): + return ToolResult(False, f"Path does not exist: {path}") + if target.is_file(): + return ToolResult(True, str(target.relative_to(self.root))) + + rows: list[str] = [] + for item in sorted(target.iterdir(), key=lambda p: (not p.is_dir(), p.name.lower())): + if item.name in SKIP_DIRS: + continue + suffix = "/" if item.is_dir() else "" + rows.append(f"{item.relative_to(self.root)}{suffix}") + if len(rows) >= 200: + rows.append("[truncated]") + break + return ToolResult(True, "\n".join(rows) or "[empty directory]") + + def read_file(self, path: str) -> ToolResult: + target = self._resolve(path) + if not target.is_file(): + return ToolResult(False, f"Not a file: {path}") + try: + text = target.read_text(encoding="utf-8") + except UnicodeDecodeError: + text = target.read_text(encoding="utf-8", errors="replace") + return ToolResult(True, text[:12000]) + + def search_text(self, query: str, path: str = ".") -> ToolResult: + if not query: + return ToolResult(False, "Missing query.") + target = self._resolve(path) + rows: list[str] = [] + files = [target] if target.is_file() else target.rglob("*") + for file_path in files: + if len(rows) >= 80: + rows.append("[truncated]") + break + if not file_path.is_file() or any(part in SKIP_DIRS for part in file_path.parts): + continue + try: + lines = file_path.read_text(encoding="utf-8", errors="replace").splitlines() + except OSError: + continue + for line_no, line in enumerate(lines, start=1): + if query.lower() in line.lower(): + rel = file_path.relative_to(self.root) + rows.append(f"{rel}:{line_no}: {line[:240]}") + if len(rows) >= 80: + break + return ToolResult(True, "\n".join(rows) or "[no matches]") + + def create_file(self, path: str, content: str) -> ToolResult: + target = self._resolve(path) + if self._is_protected(target): + return ToolResult(False, f"Refusing to write protected path: {path}") + if target.exists(): + return ToolResult(False, f"File already exists: {path}") + target.parent.mkdir(parents=True, exist_ok=True) + target.write_text(content, encoding="utf-8") + return ToolResult(True, f"Created {target.relative_to(self.root)} ({len(content)} bytes).") + + def replace_in_file(self, path: str, old: str, new: str) -> ToolResult: + if not old: + return ToolResult(False, "Missing old text.") + target = self._resolve(path) + if self._is_protected(target): + return ToolResult(False, f"Refusing to write protected path: {path}") + if not target.is_file(): + return ToolResult(False, f"Not a file: {path}") + text = target.read_text(encoding="utf-8", errors="replace") + count = text.count(old) + if count != 1: + return ToolResult(False, f"Expected exactly one match, found {count}.") + target.write_text(text.replace(old, new, 1), encoding="utf-8") + return ToolResult(True, f"Updated {target.relative_to(self.root)}.") + + def run_command(self, cmd: str, approved: bool = False) -> ToolResult: + if not cmd: + return ToolResult(False, "Missing command.") + assessment = assess_command(cmd) + if assessment.blocked: + return ToolResult(False, f"Blocked by command safety policy: {assessment.reason}") + if assessment.needs_confirmation and not (self.auto_approve_commands or approved): + return ToolResult(False, f"Command requires confirmation: {assessment.reason}") + completed = subprocess.run( + cmd, + cwd=str(self.root), + shell=True, + capture_output=True, + env=_subprocess_env(), + timeout=self.timeout_seconds, + ) + output = [] + stdout = decode_process_output(completed.stdout) + stderr = decode_process_output(completed.stderr) + if stdout: + output.append(stdout.rstrip()) + if stderr: + output.append(stderr.rstrip()) + output.append(f"[exit code {completed.returncode}]") + return ToolResult(completed.returncode == 0, "\n".join(output)) + + def apply_patch(self, patch: str) -> ToolResult: + if not patch.strip(): + return ToolResult(False, "Missing patch.") + try: + messages = apply_patch_to_root(self.root, patch, self.protected_paths) + except PatchError as exc: + return ToolResult(False, str(exc)) + return ToolResult(True, "\n".join(messages)) + + def _resolve(self, path: str) -> Path: + if not path: + raise ValueError("Missing path.") + target = (self.root / path).resolve() + if target != self.root and self.root not in target.parents: + raise ValueError(f"Path escapes workspace: {path}") + return target + + def _is_protected(self, target: Path) -> bool: + try: + rel = target.relative_to(self.root) + except ValueError: + return True + parts = rel.parts + return any(protected in parts or str(rel).startswith(protected.rstrip("/\\") + "\\") for protected in self.protected_paths) + + +def _preview_text(text: str, limit: int = 4000) -> str: + if len(text) <= limit: + return text + return text[:limit] + f"\n[truncated, {len(text) - limit} more characters]" + + +def _subprocess_env() -> dict[str, str]: + env = os.environ.copy() + env.setdefault("PYTHONIOENCODING", "utf-8") + env.setdefault("PYTHONUTF8", "1") + return env diff --git a/tests/__init__.py b/tests/__init__.py new file mode 100644 index 0000000..d377298 --- /dev/null +++ b/tests/__init__.py @@ -0,0 +1,8 @@ +import sys +from pathlib import Path + + +ROOT = Path(__file__).resolve().parents[1] +SRC = ROOT / "src" +if str(SRC) not in sys.path: + sys.path.insert(0, str(SRC)) diff --git a/tests/test_command_safety.py b/tests/test_command_safety.py new file mode 100644 index 0000000..c24a25e --- /dev/null +++ b/tests/test_command_safety.py @@ -0,0 +1,24 @@ +import unittest + +from nlprog.command_safety import assess_command + + +class CommandSafetyTests(unittest.TestCase): + def test_allows_safe_validation_command(self): + assessment = assess_command("python -m compileall src") + self.assertEqual(assessment.level, "allow") + self.assertFalse(assessment.blocked) + + def test_confirms_package_install(self): + assessment = assess_command("pip install requests") + self.assertEqual(assessment.level, "confirm") + self.assertTrue(assessment.needs_confirmation) + + def test_blocks_format(self): + assessment = assess_command("format C:") + self.assertEqual(assessment.level, "block") + self.assertTrue(assessment.blocked) + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/test_config.py b/tests/test_config.py new file mode 100644 index 0000000..8125929 --- /dev/null +++ b/tests/test_config.py @@ -0,0 +1,94 @@ +import tempfile +import unittest +import os +from pathlib import Path +from unittest.mock import patch + +import nlprog.config as config + + +class ConfigTests(unittest.TestCase): + def test_parse_config_value(self): + self.assertTrue(config.parse_config_value("require_confirmation", "true")) + self.assertEqual(config.parse_config_value("json_repair_retries", "3"), 3) + self.assertEqual(config.parse_config_value("temperature", "0.5"), 0.5) + self.assertIsNone(config.parse_config_value("base_url", "null")) + + def test_set_config_value_uses_temp_config_file(self): + with tempfile.TemporaryDirectory() as tmp: + config_file = Path(tmp) / "config.json" + with patch.object(config, "CONFIG_DIR", Path(tmp)), patch.object(config, "CONFIG_FILE", config_file): + config.write_default_config() + config.set_config_value("provider", "opencode") + data = config.read_config_file() + self.assertEqual(data["provider"], "opencode") + + def test_active_model_overrides_loaded_config(self): + with tempfile.TemporaryDirectory() as tmp: + config_file = Path(tmp) / "config.json" + with patch.object(config, "CONFIG_DIR", Path(tmp)): + with patch.object(config, "CONFIG_FILE", config_file): + with patch.dict(os.environ, {}, clear=True): + config.write_default_config() + config.add_model( + "deepseek", + provider="openai-compatible", + model="deepseek-chat", + base_url="https://api.deepseek.com/v1", + api_key_env="DEEPSEEK_API_KEY", + activate=True, + ) + + loaded = config.load_config() + + self.assertEqual(loaded.provider, "openai-compatible") + self.assertEqual(loaded.model, "deepseek-chat") + self.assertEqual(loaded.base_url, "https://api.deepseek.com/v1") + self.assertEqual(loaded.api_key_env, "DEEPSEEK_API_KEY") + + def test_legacy_config_without_active_model_still_loads_top_level_values(self): + with tempfile.TemporaryDirectory() as tmp: + config_file = Path(tmp) / "config.json" + with patch.object(config, "CONFIG_DIR", Path(tmp)): + with patch.object(config, "CONFIG_FILE", config_file): + with patch.dict(os.environ, {}, clear=True): + config.write_config_file( + { + "provider": "opencode", + "model": "", + "opencode_command": "opencode", + } + ) + + loaded = config.load_config() + active, models = config.list_models() + + self.assertEqual(loaded.provider, "opencode") + self.assertEqual(loaded.model, "") + self.assertEqual(active, "current") + self.assertEqual(models["current"]["provider"], "opencode") + + def test_model_registry_use_and_remove(self): + with tempfile.TemporaryDirectory() as tmp: + config_file = Path(tmp) / "config.json" + with patch.object(config, "CONFIG_DIR", Path(tmp)), patch.object(config, "CONFIG_FILE", config_file): + config.write_default_config() + config.add_model("codex", provider="opencode", activate=True) + config.add_model("mock2", provider="mock", model="mock-model") + + active, models = config.list_models() + self.assertEqual(active, "codex") + self.assertIn("codex", models) + self.assertIn("mock2", models) + + with self.assertRaises(ValueError): + config.remove_model("codex") + + config.use_model("mock2") + config.remove_model("codex") + _, models = config.list_models() + self.assertNotIn("codex", models) + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/test_json_repair.py b/tests/test_json_repair.py new file mode 100644 index 0000000..66ddcb3 --- /dev/null +++ b/tests/test_json_repair.py @@ -0,0 +1,38 @@ +import unittest + +from nlprog.json_repair import complete_json, parse_json_object +from nlprog.llm import Message + + +class BrokenThenGoodClient: + def __init__(self): + self.calls = 0 + + def complete(self, messages): + self.calls += 1 + if self.calls == 1: + return "not json" + return '{"final": "repaired"}' + + +class JsonRepairTests(unittest.TestCase): + def test_parse_json_object_strips_fences(self): + data, error = parse_json_object('```json\n{"ok": true}\n```') + self.assertEqual(error, "") + self.assertEqual(data, {"ok": True}) + + def test_complete_json_repairs_bad_response(self): + result = complete_json( + BrokenThenGoodClient(), + [Message("system", "test"), Message("user", "return json")], + '{"final": "summary"}', + lambda data: None if "final" in data else "missing final", + max_retries=2, + ) + self.assertEqual(result.data, {"final": "repaired"}) + self.assertEqual(result.attempts, 2) + self.assertTrue(result.repaired) + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/test_patching.py b/tests/test_patching.py new file mode 100644 index 0000000..74deaef --- /dev/null +++ b/tests/test_patching.py @@ -0,0 +1,46 @@ +import tempfile +import unittest +from pathlib import Path + +from nlprog.patching import PatchError, apply_patch_to_root, summarize_patch + + +class PatchingTests(unittest.TestCase): + def test_add_update_delete_file(self): + with tempfile.TemporaryDirectory() as tmp: + root = Path(tmp) + add_patch = """*** Begin Patch +*** Add File: notes.txt ++hello +*** End Patch""" + self.assertIn("Add notes.txt", summarize_patch(add_patch)) + self.assertEqual(apply_patch_to_root(root, add_patch, set()), ["Added notes.txt"]) + self.assertEqual((root / "notes.txt").read_text(encoding="utf-8"), "hello\n") + + update_patch = """*** Begin Patch +*** Update File: notes.txt +@@ +-hello ++hello world +*** End Patch""" + self.assertEqual(apply_patch_to_root(root, update_patch, set()), ["Updated notes.txt"]) + self.assertEqual((root / "notes.txt").read_text(encoding="utf-8"), "hello world\n") + + delete_patch = """*** Begin Patch +*** Delete File: notes.txt +*** End Patch""" + self.assertEqual(apply_patch_to_root(root, delete_patch, set()), ["Deleted notes.txt"]) + self.assertFalse((root / "notes.txt").exists()) + + def test_protected_path_is_rejected(self): + with tempfile.TemporaryDirectory() as tmp: + patch = """*** Begin Patch +*** Add File: __pycache__/blocked.txt ++x +*** End Patch""" + with self.assertRaises(PatchError): + apply_patch_to_root(Path(tmp), patch, {"__pycache__"}) + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/test_project_profile.py b/tests/test_project_profile.py new file mode 100644 index 0000000..51b961c --- /dev/null +++ b/tests/test_project_profile.py @@ -0,0 +1,29 @@ +import tempfile +import unittest +from pathlib import Path + +from nlprog.project_profile import infer_project, init_project, load_project_profile + + +class ProjectProfileTests(unittest.TestCase): + def test_infers_python_project(self): + with tempfile.TemporaryDirectory() as tmp: + root = Path(tmp) + (root / "pyproject.toml").write_text("[project]\nname='x'\n", encoding="utf-8") + (root / "src").mkdir() + profile = infer_project(root) + self.assertIn("python", profile.project_types) + self.assertIn("python -m compileall src", profile.verification_commands) + + def test_init_project_writes_profile_files(self): + with tempfile.TemporaryDirectory() as tmp: + root = Path(tmp) + (root / "pyproject.toml").write_text("[project]\nname='x'\n", encoding="utf-8") + created = init_project(root) + self.assertEqual(len(created), 3) + profile = load_project_profile(root) + self.assertIn("pyproject.toml", profile.important_files) + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/test_run_log.py b/tests/test_run_log.py new file mode 100644 index 0000000..4b8767c --- /dev/null +++ b/tests/test_run_log.py @@ -0,0 +1,29 @@ +import tempfile +import unittest +from pathlib import Path + +from nlprog.run_log import list_run_logs, load_run_log, write_run_log + + +class Event: + def __init__(self, kind, message): + self.kind = kind + self.message = message + + +class RunLogTests(unittest.TestCase): + def test_write_list_and_load_latest_log(self): + with tempfile.TemporaryDirectory() as tmp: + root = Path(tmp) + path = write_run_log(root, "task", "done", [Event("final", "done")], "2026-01-01T00:00:00") + self.assertTrue(path.exists()) + logs = list_run_logs(root) + self.assertEqual(logs, [path]) + loaded_path, data = load_run_log(root, "latest") + self.assertEqual(loaded_path, path) + self.assertEqual(data["task"], "task") + self.assertEqual(data["final"], "done") + + +if __name__ == "__main__": + unittest.main()