mirror of
https://github.com/basicmachines-co/basic-memory
synced 2026-06-21 13:47:35 +00:00
chore: apply lint and formatting fixes for 0.14.4 release (#290)
Signed-off-by: Joe P <joe@basicmemory.com> Co-authored-by: Claude <noreply@anthropic.com>
This commit is contained in:
+49
-42
@@ -7,82 +7,89 @@ import json
|
||||
import csv
|
||||
from pathlib import Path
|
||||
|
||||
|
||||
def create_csv_exhibits():
|
||||
"""Create CSV Exhibit A files for each contributor."""
|
||||
|
||||
|
||||
# Read the JSON inventory
|
||||
inventory_files = list(Path("legal_inventory_main").glob("*.json"))
|
||||
if not inventory_files:
|
||||
print("Error: No JSON inventory files found")
|
||||
return
|
||||
|
||||
|
||||
inventory_file = inventory_files[0] # Use the most recent one
|
||||
print(f"Using inventory file: {inventory_file}")
|
||||
|
||||
with open(inventory_file, 'r') as f:
|
||||
|
||||
with open(inventory_file, "r") as f:
|
||||
data = json.load(f)
|
||||
|
||||
files = data['files']
|
||||
|
||||
|
||||
files = data["files"]
|
||||
|
||||
# Create output directory
|
||||
output_dir = Path("legal_exhibits")
|
||||
output_dir.mkdir(exist_ok=True)
|
||||
|
||||
|
||||
# Contributors we need exhibits for
|
||||
target_contributors = {
|
||||
'jope-bm': 'joe_exhibit_a.csv',
|
||||
'Drew Cain': 'drew_cain_exhibit_a.csv'
|
||||
}
|
||||
|
||||
target_contributors = {"jope-bm": "joe_exhibit_a.csv", "Drew Cain": "drew_cain_exhibit_a.csv"}
|
||||
|
||||
print("Creating CSV exhibits for contributors...")
|
||||
|
||||
|
||||
for contributor_key, filename in target_contributors.items():
|
||||
# Find files for this contributor
|
||||
contributor_files = []
|
||||
|
||||
|
||||
for file_info in files:
|
||||
# Check if this contributor is listed in the file's contributors
|
||||
for contrib in file_info.get('contributors', []):
|
||||
if contributor_key in contrib['name']:
|
||||
for contrib in file_info.get("contributors", []):
|
||||
if contributor_key in contrib["name"]:
|
||||
contributor_files.append(file_info)
|
||||
break
|
||||
|
||||
|
||||
if not contributor_files:
|
||||
print(f"No files found for {contributor_key}")
|
||||
continue
|
||||
|
||||
|
||||
# Sort files by path
|
||||
contributor_files.sort(key=lambda x: x['path'])
|
||||
|
||||
contributor_files.sort(key=lambda x: x["path"])
|
||||
|
||||
# Create CSV file
|
||||
csv_file = output_dir / filename
|
||||
|
||||
with open(csv_file, 'w', newline='', encoding='utf-8') as csvfile:
|
||||
|
||||
with open(csv_file, "w", newline="", encoding="utf-8") as csvfile:
|
||||
fieldnames = [
|
||||
'file_path', 'file_name', 'category', 'size_bytes',
|
||||
'modified_date', 'primary_author', 'all_contributors', 'sha256_hash'
|
||||
"file_path",
|
||||
"file_name",
|
||||
"category",
|
||||
"size_bytes",
|
||||
"modified_date",
|
||||
"primary_author",
|
||||
"all_contributors",
|
||||
"sha256_hash",
|
||||
]
|
||||
|
||||
|
||||
writer = csv.DictWriter(csvfile, fieldnames=fieldnames)
|
||||
writer.writeheader()
|
||||
|
||||
|
||||
for file_info in contributor_files:
|
||||
contributors_list = '; '.join([c['name'] for c in file_info['contributors']])
|
||||
|
||||
writer.writerow({
|
||||
'file_path': file_info['path'],
|
||||
'file_name': file_info['name'],
|
||||
'category': file_info['category'],
|
||||
'size_bytes': file_info['size_bytes'],
|
||||
'modified_date': file_info['modified_time'][:10],
|
||||
'primary_author': file_info['primary_author'],
|
||||
'all_contributors': contributors_list,
|
||||
'sha256_hash': file_info['sha256_hash']
|
||||
})
|
||||
|
||||
contributors_list = "; ".join([c["name"] for c in file_info["contributors"]])
|
||||
|
||||
writer.writerow(
|
||||
{
|
||||
"file_path": file_info["path"],
|
||||
"file_name": file_info["name"],
|
||||
"category": file_info["category"],
|
||||
"size_bytes": file_info["size_bytes"],
|
||||
"modified_date": file_info["modified_time"][:10],
|
||||
"primary_author": file_info["primary_author"],
|
||||
"all_contributors": contributors_list,
|
||||
"sha256_hash": file_info["sha256_hash"],
|
||||
}
|
||||
)
|
||||
|
||||
print(f"Created CSV exhibit for {contributor_key}: {csv_file}")
|
||||
print(f" - {len(contributor_files)} files")
|
||||
print(f" - {sum(f['size_bytes'] for f in contributor_files):,} bytes")
|
||||
|
||||
if __name__ == '__main__':
|
||||
create_csv_exhibits()
|
||||
|
||||
if __name__ == "__main__":
|
||||
create_csv_exhibits()
|
||||
|
||||
@@ -7,51 +7,52 @@ import json
|
||||
from pathlib import Path
|
||||
from datetime import datetime
|
||||
|
||||
|
||||
def create_individual_exhibits():
|
||||
"""Create individual Exhibit A files for each contributor."""
|
||||
|
||||
|
||||
# Read the JSON inventory
|
||||
inventory_file = Path("legal_inventory_main/basic_memory_inventory_20250730_101521.json")
|
||||
|
||||
|
||||
if not inventory_file.exists():
|
||||
print(f"Error: {inventory_file} not found")
|
||||
return
|
||||
|
||||
with open(inventory_file, 'r') as f:
|
||||
|
||||
with open(inventory_file, "r") as f:
|
||||
data = json.load(f)
|
||||
|
||||
files = data['files']
|
||||
|
||||
|
||||
files = data["files"]
|
||||
|
||||
# Create output directory
|
||||
output_dir = Path("legal_exhibits")
|
||||
output_dir.mkdir(exist_ok=True)
|
||||
|
||||
|
||||
# Contributors we need exhibits for (based on copyright assignments)
|
||||
target_contributors = {
|
||||
'jope-bm': 'Joseph "Joe" [Last Name]', # Need to get his full name
|
||||
'Drew Cain': 'Drew Cain'
|
||||
"jope-bm": 'Joseph "Joe" [Last Name]', # Need to get his full name
|
||||
"Drew Cain": "Drew Cain",
|
||||
}
|
||||
|
||||
|
||||
print("Creating individual contributor exhibits...")
|
||||
|
||||
|
||||
for contributor_key, full_name in target_contributors.items():
|
||||
# Find files for this contributor
|
||||
contributor_files = []
|
||||
|
||||
|
||||
for file_info in files:
|
||||
# Check if this contributor is listed in the file's contributors
|
||||
for contrib in file_info.get('contributors', []):
|
||||
if contributor_key in contrib['name']:
|
||||
for contrib in file_info.get("contributors", []):
|
||||
if contributor_key in contrib["name"]:
|
||||
contributor_files.append(file_info)
|
||||
break
|
||||
|
||||
|
||||
if not contributor_files:
|
||||
print(f"No files found for {contributor_key}")
|
||||
continue
|
||||
|
||||
|
||||
# Sort files by path
|
||||
contributor_files.sort(key=lambda x: x['path'])
|
||||
|
||||
contributor_files.sort(key=lambda x: x["path"])
|
||||
|
||||
# Create exhibit markdown
|
||||
exhibit_content = f"""# Exhibit A - Assigned Works
|
||||
## Copyright Assignment: {full_name} to Basic Memory LLC
|
||||
@@ -62,44 +63,44 @@ def create_individual_exhibits():
|
||||
|
||||
## Summary
|
||||
- **Total Files:** {len(contributor_files)}
|
||||
- **Total Size:** {sum(f['size_bytes'] for f in contributor_files):,} bytes
|
||||
- **Categories:** {', '.join(set(f['category'] for f in contributor_files))}
|
||||
- **Total Size:** {sum(f["size_bytes"] for f in contributor_files):,} bytes
|
||||
- **Categories:** {", ".join(set(f["category"] for f in contributor_files))}
|
||||
|
||||
## Detailed File List
|
||||
|
||||
"""
|
||||
|
||||
|
||||
# Group by category
|
||||
categories = {}
|
||||
for file_info in contributor_files:
|
||||
category = file_info['category']
|
||||
category = file_info["category"]
|
||||
if category not in categories:
|
||||
categories[category] = []
|
||||
categories[category].append(file_info)
|
||||
|
||||
|
||||
# Add files by category
|
||||
for category, category_files in sorted(categories.items()):
|
||||
exhibit_content += f"### {category.replace('_', ' ').title()}\n\n"
|
||||
|
||||
|
||||
for file_info in category_files:
|
||||
exhibit_content += f"**{file_info['path']}**\n"
|
||||
exhibit_content += f"- Size: {file_info['size_bytes']:,} bytes\n"
|
||||
exhibit_content += f"- Modified: {file_info['modified_time'][:10]}\n"
|
||||
exhibit_content += f"- Primary Author: {file_info['primary_author']}\n"
|
||||
|
||||
|
||||
# Show all contributors for this file
|
||||
if len(file_info['contributors']) > 1:
|
||||
contributors_list = ', '.join([c['name'] for c in file_info['contributors']])
|
||||
if len(file_info["contributors"]) > 1:
|
||||
contributors_list = ", ".join([c["name"] for c in file_info["contributors"]])
|
||||
exhibit_content += f"- All Contributors: {contributors_list}\n"
|
||||
|
||||
|
||||
exhibit_content += f"- SHA-256: `{file_info['sha256_hash']}`\n\n"
|
||||
|
||||
|
||||
# Add verification section
|
||||
exhibit_content += f"""
|
||||
## Verification
|
||||
This exhibit lists all files in the Basic Memory repository where {full_name} is identified as a contributor based on git commit history analysis.
|
||||
|
||||
**Analysis Date:** {datetime.now().strftime('%Y-%m-%d')}
|
||||
**Analysis Date:** {datetime.now().strftime("%Y-%m-%d")}
|
||||
**Repository State:** Basic Memory main branch
|
||||
**Method:** Git history analysis via `git log --follow` for each file
|
||||
|
||||
@@ -110,54 +111,55 @@ This exhibit lists all files in the Basic Memory repository where {full_name} is
|
||||
|
||||
*This exhibit is attached to and forms part of the Copyright Assignment Agreement between {full_name} and Basic Memory LLC.*
|
||||
"""
|
||||
|
||||
|
||||
# Write exhibit file
|
||||
safe_name = contributor_key.replace(' ', '_').replace('-', '_').lower()
|
||||
safe_name = contributor_key.replace(" ", "_").replace("-", "_").lower()
|
||||
exhibit_file = output_dir / f"exhibit_a_{safe_name}.md"
|
||||
|
||||
with open(exhibit_file, 'w') as f:
|
||||
|
||||
with open(exhibit_file, "w") as f:
|
||||
f.write(exhibit_content)
|
||||
|
||||
|
||||
print(f"Created exhibit for {full_name}: {exhibit_file}")
|
||||
print(f" - {len(contributor_files)} files")
|
||||
print(f" - {sum(f['size_bytes'] for f in contributor_files):,} bytes")
|
||||
|
||||
|
||||
# Create overall summary exhibit (for Paul's assignment to Basic Machines LLC)
|
||||
create_overall_summary_exhibit(data, output_dir)
|
||||
|
||||
|
||||
def create_overall_summary_exhibit(data, output_dir):
|
||||
"""Create overall summary exhibit for Company Agreement."""
|
||||
|
||||
files = data['files']
|
||||
summary = data['summary']
|
||||
contributors = data['contributors']
|
||||
|
||||
|
||||
files = data["files"]
|
||||
summary = data["summary"]
|
||||
contributors = data["contributors"]
|
||||
|
||||
summary_content = f"""# Basic Memory Repository - Complete IP Inventory
|
||||
## For Basic Memory LLC Company Agreement
|
||||
|
||||
**Analysis Date:** {summary['scan_date'][:10]}
|
||||
**Repository:** {summary['repository_path']}
|
||||
**Analysis Date:** {summary["scan_date"][:10]}
|
||||
**Repository:** {summary["repository_path"]}
|
||||
|
||||
## Executive Summary
|
||||
- **Total Files:** {summary['total_files']:,}
|
||||
- **Total Size:** {summary['total_size_bytes']:,} bytes
|
||||
- **Contributors:** {summary['contributor_count']}
|
||||
- **Primary Author:** Paul Hernandez ({len(contributors.get('Paul Hernandez', {}).get('files', []))} files)
|
||||
- **Total Files:** {summary["total_files"]:,}
|
||||
- **Total Size:** {summary["total_size_bytes"]:,} bytes
|
||||
- **Contributors:** {summary["contributor_count"]}
|
||||
- **Primary Author:** Paul Hernandez ({len(contributors.get("Paul Hernandez", {}).get("files", []))} files)
|
||||
|
||||
## File Categories
|
||||
"""
|
||||
|
||||
for category, count in sorted(summary['categories'].items()):
|
||||
|
||||
for category, count in sorted(summary["categories"].items()):
|
||||
summary_content += f"- **{category.replace('_', ' ').title()}:** {count} files\n"
|
||||
|
||||
|
||||
summary_content += """
|
||||
|
||||
## Contributor Summary
|
||||
"""
|
||||
|
||||
for contrib in summary['top_contributors']:
|
||||
|
||||
for contrib in summary["top_contributors"]:
|
||||
summary_content += f"- **{contrib['name']}** ({contrib['email']}): {contrib['file_count']} files, {contrib['commit_count']} commits\n"
|
||||
|
||||
|
||||
summary_content += """
|
||||
|
||||
## Legal Significance
|
||||
@@ -177,32 +179,33 @@ This comprehensive file inventory serves as:
|
||||
## Repository Contents by Category
|
||||
|
||||
"""
|
||||
|
||||
|
||||
# Add sample files by category (first 10 in each category)
|
||||
for category in sorted(summary['categories'].keys()):
|
||||
category_files = [f for f in files if f['category'] == category][:10]
|
||||
for category in sorted(summary["categories"].keys()):
|
||||
category_files = [f for f in files if f["category"] == category][:10]
|
||||
if category_files:
|
||||
summary_content += f"### {category.replace('_', ' ').title()} (Sample)\n\n"
|
||||
for file_info in category_files:
|
||||
summary_content += f"- `{file_info['path']}` ({file_info['size_bytes']:,} bytes)\n"
|
||||
|
||||
if len([f for f in files if f['category'] == category]) > 10:
|
||||
remaining = len([f for f in files if f['category'] == category]) - 10
|
||||
|
||||
if len([f for f in files if f["category"] == category]) > 10:
|
||||
remaining = len([f for f in files if f["category"] == category]) - 10
|
||||
summary_content += f"- *... and {remaining} more files*\n"
|
||||
summary_content += "\n"
|
||||
|
||||
|
||||
summary_content += f"""
|
||||
---
|
||||
|
||||
*This inventory was generated automatically from git repository analysis and represents the complete intellectual property foundation of Basic Memory as of {summary['scan_date'][:10]}.*
|
||||
*This inventory was generated automatically from git repository analysis and represents the complete intellectual property foundation of Basic Memory as of {summary["scan_date"][:10]}.*
|
||||
"""
|
||||
|
||||
|
||||
# Write summary file
|
||||
summary_file = output_dir / "basic_memory_complete_inventory.md"
|
||||
with open(summary_file, 'w') as f:
|
||||
with open(summary_file, "w") as f:
|
||||
f.write(summary_content)
|
||||
|
||||
|
||||
print(f"Created complete inventory summary: {summary_file}")
|
||||
|
||||
if __name__ == '__main__':
|
||||
create_individual_exhibits()
|
||||
|
||||
if __name__ == "__main__":
|
||||
create_individual_exhibits()
|
||||
|
||||
+236
-188
@@ -25,65 +25,110 @@ from typing import Dict, List
|
||||
import argparse
|
||||
import hashlib
|
||||
|
||||
|
||||
class FileInventoryGenerator:
|
||||
def __init__(self, repo_path: str = "."):
|
||||
self.repo_path = Path(repo_path).resolve()
|
||||
self.inventory = []
|
||||
|
||||
|
||||
# File patterns to exclude from legal inventory
|
||||
self.exclude_patterns = {
|
||||
# Version control and git
|
||||
'.git', '.gitignore', '.gitmodules',
|
||||
|
||||
".git",
|
||||
".gitignore",
|
||||
".gitmodules",
|
||||
# Python compiled and cache files
|
||||
'__pycache__', '*.pyc', '*.pyo', '*.pyd', '.pytest_cache',
|
||||
|
||||
"__pycache__",
|
||||
"*.pyc",
|
||||
"*.pyo",
|
||||
"*.pyd",
|
||||
".pytest_cache",
|
||||
# Virtual environments and dependencies
|
||||
'.venv', 'venv', '.env', 'env', 'ENV',
|
||||
'*.dist-info', 'site-packages',
|
||||
|
||||
".venv",
|
||||
"venv",
|
||||
".env",
|
||||
"env",
|
||||
"ENV",
|
||||
"*.dist-info",
|
||||
"site-packages",
|
||||
# IDE and editor files
|
||||
'.idea', '.vscode', '*.swp', '*.swo', '.DS_Store',
|
||||
|
||||
".idea",
|
||||
".vscode",
|
||||
"*.swp",
|
||||
"*.swo",
|
||||
".DS_Store",
|
||||
# Build and distribution artifacts
|
||||
'build', 'dist', 'htmlcov', '.coverage', '.coverage.*',
|
||||
'*.egg-info', '.eggs', 'wheels',
|
||||
|
||||
"build",
|
||||
"dist",
|
||||
"htmlcov",
|
||||
".coverage",
|
||||
".coverage.*",
|
||||
"*.egg-info",
|
||||
".eggs",
|
||||
"wheels",
|
||||
# Cache and temporary files
|
||||
'.ruff_cache', '.mypy_cache', '.tox',
|
||||
'node_modules', '.npm',
|
||||
|
||||
".ruff_cache",
|
||||
".mypy_cache",
|
||||
".tox",
|
||||
"node_modules",
|
||||
".npm",
|
||||
# Documentation build artifacts (but keep source docs)
|
||||
'.obsidian',
|
||||
|
||||
".obsidian",
|
||||
# Lock files (these are generated)
|
||||
'uv.lock', 'Pipfile.lock', 'poetry.lock', 'package-lock.json'
|
||||
"uv.lock",
|
||||
"Pipfile.lock",
|
||||
"poetry.lock",
|
||||
"package-lock.json",
|
||||
}
|
||||
|
||||
|
||||
# File extensions that are definitely source/authored content
|
||||
self.source_extensions = {
|
||||
'.py', '.md', '.rst', '.txt', '.toml', '.yaml', '.yml',
|
||||
'.json', '.cfg', '.ini', '.conf', '.sh', '.sql',
|
||||
'.js', '.ts', '.jsx', '.tsx', '.css', '.scss', '.sass',
|
||||
'.html', '.htm', '.xml', '.svg', '.dockerfile', '.Dockerfile'
|
||||
".py",
|
||||
".md",
|
||||
".rst",
|
||||
".txt",
|
||||
".toml",
|
||||
".yaml",
|
||||
".yml",
|
||||
".json",
|
||||
".cfg",
|
||||
".ini",
|
||||
".conf",
|
||||
".sh",
|
||||
".sql",
|
||||
".js",
|
||||
".ts",
|
||||
".jsx",
|
||||
".tsx",
|
||||
".css",
|
||||
".scss",
|
||||
".sass",
|
||||
".html",
|
||||
".htm",
|
||||
".xml",
|
||||
".svg",
|
||||
".dockerfile",
|
||||
".Dockerfile",
|
||||
}
|
||||
|
||||
|
||||
# License file patterns
|
||||
self.license_patterns = {
|
||||
'LICENSE', 'LICENCE', 'COPYING', 'COPYRIGHT',
|
||||
'license.txt', 'LICENSE.txt', 'LICENSE.md',
|
||||
'CITATION.cff', 'CLA.md'
|
||||
"LICENSE",
|
||||
"LICENCE",
|
||||
"COPYING",
|
||||
"COPYRIGHT",
|
||||
"license.txt",
|
||||
"LICENSE.txt",
|
||||
"LICENSE.md",
|
||||
"CITATION.cff",
|
||||
"CLA.md",
|
||||
}
|
||||
|
||||
def run_git_command(self, command: List[str]) -> str:
|
||||
"""Run a git command and return the output."""
|
||||
try:
|
||||
result = subprocess.run(
|
||||
['git'] + command,
|
||||
cwd=self.repo_path,
|
||||
capture_output=True,
|
||||
text=True,
|
||||
check=True
|
||||
["git"] + command, cwd=self.repo_path, capture_output=True, text=True, check=True
|
||||
)
|
||||
return result.stdout.strip()
|
||||
except subprocess.CalledProcessError:
|
||||
@@ -92,14 +137,14 @@ class FileInventoryGenerator:
|
||||
def get_file_contributors(self, file_path: str) -> Dict[str, int]:
|
||||
"""Get contributors and their line contributions for a file."""
|
||||
try:
|
||||
blame_output = self.run_git_command(['blame', '--line-porcelain', file_path])
|
||||
blame_output = self.run_git_command(["blame", "--line-porcelain", file_path])
|
||||
contributors = {}
|
||||
|
||||
for line in blame_output.split('\n'):
|
||||
if line.startswith('author '):
|
||||
|
||||
for line in blame_output.split("\n"):
|
||||
if line.startswith("author "):
|
||||
author = line[7:] # Remove 'author ' prefix
|
||||
contributors[author] = contributors.get(author, 0) + 1
|
||||
|
||||
|
||||
return contributors
|
||||
except Exception:
|
||||
return {}
|
||||
@@ -108,61 +153,59 @@ class FileInventoryGenerator:
|
||||
"""Get file creation date, last modification, and total commits."""
|
||||
try:
|
||||
# Get creation date (first commit)
|
||||
first_commit = self.run_git_command([
|
||||
'log', '--follow', '--format=%ad', '--date=iso',
|
||||
'--reverse', file_path
|
||||
]).split('\n')[0] if self.run_git_command([
|
||||
'log', '--follow', '--format=%ad', '--date=iso',
|
||||
'--reverse', file_path
|
||||
]) else None
|
||||
|
||||
first_commit = (
|
||||
self.run_git_command(
|
||||
["log", "--follow", "--format=%ad", "--date=iso", "--reverse", file_path]
|
||||
).split("\n")[0]
|
||||
if self.run_git_command(
|
||||
["log", "--follow", "--format=%ad", "--date=iso", "--reverse", file_path]
|
||||
)
|
||||
else None
|
||||
)
|
||||
|
||||
# Get last modification date
|
||||
last_commit = self.run_git_command([
|
||||
'log', '-1', '--format=%ad', '--date=iso', file_path
|
||||
])
|
||||
|
||||
last_commit = self.run_git_command(
|
||||
["log", "-1", "--format=%ad", "--date=iso", file_path]
|
||||
)
|
||||
|
||||
# Get total commits for this file
|
||||
commit_count = len(self.run_git_command([
|
||||
'log', '--follow', '--oneline', file_path
|
||||
]).split('\n')) if self.run_git_command([
|
||||
'log', '--follow', '--oneline', file_path
|
||||
]) else 0
|
||||
|
||||
commit_count = (
|
||||
len(self.run_git_command(["log", "--follow", "--oneline", file_path]).split("\n"))
|
||||
if self.run_git_command(["log", "--follow", "--oneline", file_path])
|
||||
else 0
|
||||
)
|
||||
|
||||
return {
|
||||
'created': first_commit or 'Unknown',
|
||||
'last_modified': last_commit or 'Unknown',
|
||||
'commit_count': commit_count
|
||||
"created": first_commit or "Unknown",
|
||||
"last_modified": last_commit or "Unknown",
|
||||
"commit_count": commit_count,
|
||||
}
|
||||
except Exception:
|
||||
return {
|
||||
'created': 'Unknown',
|
||||
'last_modified': 'Unknown',
|
||||
'commit_count': 0
|
||||
}
|
||||
return {"created": "Unknown", "last_modified": "Unknown", "commit_count": 0}
|
||||
|
||||
def should_exclude_file(self, file_path: Path) -> bool:
|
||||
"""Determine if a file should be excluded from the inventory."""
|
||||
str_path = str(file_path)
|
||||
|
||||
|
||||
# Check if any part of the path matches exclude patterns
|
||||
for pattern in self.exclude_patterns:
|
||||
if pattern in str_path or file_path.match(pattern):
|
||||
return True
|
||||
|
||||
|
||||
# Exclude files in virtual environment paths
|
||||
if '/.venv/' in str_path or '/venv/' in str_path:
|
||||
if "/.venv/" in str_path or "/venv/" in str_path:
|
||||
return True
|
||||
|
||||
|
||||
# Exclude binary files that are likely dependencies
|
||||
if file_path.suffix in {'.so', '.dylib', '.dll', '.pyd'}:
|
||||
if file_path.suffix in {".so", ".dylib", ".dll", ".pyd"}:
|
||||
return True
|
||||
|
||||
|
||||
return False
|
||||
|
||||
def calculate_file_hash(self, file_path: Path) -> str:
|
||||
"""Calculate SHA-256 hash of file content."""
|
||||
try:
|
||||
with open(file_path, 'rb') as f:
|
||||
with open(file_path, "rb") as f:
|
||||
return hashlib.sha256(f.read()).hexdigest()
|
||||
except Exception:
|
||||
return ""
|
||||
@@ -170,59 +213,65 @@ class FileInventoryGenerator:
|
||||
def categorize_file(self, file_path: Path) -> str:
|
||||
"""Categorize the file based on its path and extension."""
|
||||
str_path = str(file_path).lower()
|
||||
|
||||
|
||||
# License and legal files
|
||||
if any(pattern.lower() in file_path.name.lower() for pattern in self.license_patterns):
|
||||
return "Legal/License"
|
||||
|
||||
|
||||
# Documentation
|
||||
if file_path.suffix.lower() in {'.md', '.rst', '.txt'} and any(
|
||||
doc_dir in str_path for doc_dir in ['doc', 'readme', 'changelog', 'contributing']
|
||||
if file_path.suffix.lower() in {".md", ".rst", ".txt"} and any(
|
||||
doc_dir in str_path for doc_dir in ["doc", "readme", "changelog", "contributing"]
|
||||
):
|
||||
return "Documentation"
|
||||
|
||||
|
||||
# Configuration files
|
||||
if file_path.suffix.lower() in {'.toml', '.yaml', '.yml', '.json', '.cfg', '.ini', '.conf'}:
|
||||
if file_path.suffix.lower() in {".toml", ".yaml", ".yml", ".json", ".cfg", ".ini", ".conf"}:
|
||||
return "Configuration"
|
||||
|
||||
|
||||
# Source code
|
||||
if file_path.suffix.lower() in {'.py', '.js', '.ts', '.jsx', '.tsx'}:
|
||||
if file_path.suffix.lower() in {".py", ".js", ".ts", ".jsx", ".tsx"}:
|
||||
return "Source Code"
|
||||
|
||||
|
||||
# Tests
|
||||
if 'test' in str_path and file_path.suffix.lower() == '.py':
|
||||
if "test" in str_path and file_path.suffix.lower() == ".py":
|
||||
return "Test Code"
|
||||
|
||||
|
||||
# Build and deployment
|
||||
if file_path.name.lower() in {'dockerfile', 'justfile', 'makefile'} or file_path.suffix.lower() in {'.sh'}:
|
||||
if file_path.name.lower() in {
|
||||
"dockerfile",
|
||||
"justfile",
|
||||
"makefile",
|
||||
} or file_path.suffix.lower() in {".sh"}:
|
||||
return "Build/Deployment"
|
||||
|
||||
|
||||
# Database and migrations
|
||||
if 'migration' in str_path or 'alembic' in str_path or file_path.suffix.lower() == '.sql':
|
||||
if "migration" in str_path or "alembic" in str_path or file_path.suffix.lower() == ".sql":
|
||||
return "Database/Migration"
|
||||
|
||||
|
||||
# Templates and resources
|
||||
if file_path.suffix.lower() in {'.hbs', '.j2', '.jinja', '.template'}:
|
||||
if file_path.suffix.lower() in {".hbs", ".j2", ".jinja", ".template"}:
|
||||
return "Templates/Resources"
|
||||
|
||||
|
||||
return "Other"
|
||||
|
||||
def scan_repository(self):
|
||||
"""Scan the repository and build the file inventory."""
|
||||
print(f"Scanning repository: {self.repo_path}")
|
||||
|
||||
|
||||
for root, dirs, files in os.walk(self.repo_path):
|
||||
# Skip excluded directories
|
||||
dirs[:] = [d for d in dirs if not any(pattern in d for pattern in self.exclude_patterns)]
|
||||
|
||||
dirs[:] = [
|
||||
d for d in dirs if not any(pattern in d for pattern in self.exclude_patterns)
|
||||
]
|
||||
|
||||
for file in files:
|
||||
file_path = Path(root) / file
|
||||
relative_path = file_path.relative_to(self.repo_path)
|
||||
|
||||
|
||||
# Skip excluded files
|
||||
if self.should_exclude_file(relative_path):
|
||||
continue
|
||||
|
||||
|
||||
# Get file stats
|
||||
try:
|
||||
stat_info = file_path.stat()
|
||||
@@ -231,190 +280,190 @@ class FileInventoryGenerator:
|
||||
except Exception:
|
||||
file_size = 0
|
||||
modified_time = datetime.now()
|
||||
|
||||
|
||||
# Get git information
|
||||
contributors = self.get_file_contributors(str(relative_path))
|
||||
history = self.get_file_history(str(relative_path))
|
||||
|
||||
|
||||
# Calculate file hash for integrity verification
|
||||
file_hash = self.calculate_file_hash(file_path)
|
||||
|
||||
|
||||
# Build inventory entry
|
||||
entry = {
|
||||
'file_path': str(relative_path),
|
||||
'full_path': str(file_path),
|
||||
'file_name': file_path.name,
|
||||
'file_extension': file_path.suffix,
|
||||
'file_size_bytes': file_size,
|
||||
'category': self.categorize_file(relative_path),
|
||||
'fs_modified_date': modified_time.isoformat(),
|
||||
'git_created_date': history['created'],
|
||||
'git_last_modified': history['last_modified'],
|
||||
'git_commit_count': history['commit_count'],
|
||||
'contributors': contributors,
|
||||
'primary_author': max(contributors.items(), key=lambda x: x[1])[0] if contributors else 'Unknown',
|
||||
'contributor_count': len(contributors),
|
||||
'total_author_lines': sum(contributors.values()) if contributors else 0,
|
||||
'sha256_hash': file_hash,
|
||||
'scan_timestamp': datetime.now().isoformat()
|
||||
"file_path": str(relative_path),
|
||||
"full_path": str(file_path),
|
||||
"file_name": file_path.name,
|
||||
"file_extension": file_path.suffix,
|
||||
"file_size_bytes": file_size,
|
||||
"category": self.categorize_file(relative_path),
|
||||
"fs_modified_date": modified_time.isoformat(),
|
||||
"git_created_date": history["created"],
|
||||
"git_last_modified": history["last_modified"],
|
||||
"git_commit_count": history["commit_count"],
|
||||
"contributors": contributors,
|
||||
"primary_author": max(contributors.items(), key=lambda x: x[1])[0]
|
||||
if contributors
|
||||
else "Unknown",
|
||||
"contributor_count": len(contributors),
|
||||
"total_author_lines": sum(contributors.values()) if contributors else 0,
|
||||
"sha256_hash": file_hash,
|
||||
"scan_timestamp": datetime.now().isoformat(),
|
||||
}
|
||||
|
||||
|
||||
self.inventory.append(entry)
|
||||
|
||||
|
||||
print(f"Scanned {len(self.inventory)} files")
|
||||
|
||||
def get_summary_statistics(self) -> Dict:
|
||||
"""Generate summary statistics for the inventory."""
|
||||
if not self.inventory:
|
||||
return {}
|
||||
|
||||
|
||||
# Collect all contributors
|
||||
all_contributors = set()
|
||||
for entry in self.inventory:
|
||||
all_contributors.update(entry['contributors'].keys())
|
||||
|
||||
all_contributors.update(entry["contributors"].keys())
|
||||
|
||||
# Category breakdown
|
||||
categories = {}
|
||||
for entry in self.inventory:
|
||||
cat = entry['category']
|
||||
cat = entry["category"]
|
||||
categories[cat] = categories.get(cat, 0) + 1
|
||||
|
||||
|
||||
# File extension breakdown
|
||||
extensions = {}
|
||||
for entry in self.inventory:
|
||||
ext = entry['file_extension'] or 'no_extension'
|
||||
ext = entry["file_extension"] or "no_extension"
|
||||
extensions[ext] = extensions.get(ext, 0) + 1
|
||||
|
||||
|
||||
# Contributor statistics
|
||||
contributor_files = {}
|
||||
contributor_lines = {}
|
||||
for entry in self.inventory:
|
||||
for contributor, lines in entry['contributors'].items():
|
||||
for contributor, lines in entry["contributors"].items():
|
||||
contributor_files[contributor] = contributor_files.get(contributor, 0) + 1
|
||||
contributor_lines[contributor] = contributor_lines.get(contributor, 0) + lines
|
||||
|
||||
|
||||
return {
|
||||
'total_files': len(self.inventory),
|
||||
'total_contributors': len(all_contributors),
|
||||
'categories': categories,
|
||||
'file_extensions': extensions,
|
||||
'contributor_file_counts': contributor_files,
|
||||
'contributor_line_counts': contributor_lines,
|
||||
'scan_date': datetime.now().isoformat(),
|
||||
'repository_path': str(self.repo_path)
|
||||
"total_files": len(self.inventory),
|
||||
"total_contributors": len(all_contributors),
|
||||
"categories": categories,
|
||||
"file_extensions": extensions,
|
||||
"contributor_file_counts": contributor_files,
|
||||
"contributor_line_counts": contributor_lines,
|
||||
"scan_date": datetime.now().isoformat(),
|
||||
"repository_path": str(self.repo_path),
|
||||
}
|
||||
|
||||
def export_csv(self, output_file: str):
|
||||
"""Export inventory to CSV format."""
|
||||
with open(output_file, 'w', newline='', encoding='utf-8') as f:
|
||||
with open(output_file, "w", newline="", encoding="utf-8") as f:
|
||||
if not self.inventory:
|
||||
return
|
||||
|
||||
|
||||
fieldnames = list(self.inventory[0].keys())
|
||||
# Convert complex fields to strings for CSV
|
||||
fieldnames = [f for f in fieldnames if f != 'contributors']
|
||||
fieldnames.append('contributors_json')
|
||||
|
||||
fieldnames = [f for f in fieldnames if f != "contributors"]
|
||||
fieldnames.append("contributors_json")
|
||||
|
||||
writer = csv.DictWriter(f, fieldnames=fieldnames)
|
||||
writer.writeheader()
|
||||
|
||||
|
||||
for entry in self.inventory:
|
||||
row = {k: v for k, v in entry.items() if k != 'contributors'}
|
||||
row['contributors_json'] = json.dumps(entry['contributors'])
|
||||
row = {k: v for k, v in entry.items() if k != "contributors"}
|
||||
row["contributors_json"] = json.dumps(entry["contributors"])
|
||||
writer.writerow(row)
|
||||
|
||||
def export_json(self, output_file: str):
|
||||
"""Export inventory to JSON format."""
|
||||
export_data = {
|
||||
'metadata': self.get_summary_statistics(),
|
||||
'files': self.inventory
|
||||
}
|
||||
|
||||
with open(output_file, 'w', encoding='utf-8') as f:
|
||||
export_data = {"metadata": self.get_summary_statistics(), "files": self.inventory}
|
||||
|
||||
with open(output_file, "w", encoding="utf-8") as f:
|
||||
json.dump(export_data, f, indent=2, ensure_ascii=False)
|
||||
|
||||
def export_markdown(self, output_file: str):
|
||||
"""Export inventory to Markdown format."""
|
||||
stats = self.get_summary_statistics()
|
||||
|
||||
with open(output_file, 'w', encoding='utf-8') as f:
|
||||
|
||||
with open(output_file, "w", encoding="utf-8") as f:
|
||||
f.write("# Basic Memory - Legal File Inventory\n\n")
|
||||
f.write(f"**Generated:** {datetime.now().isoformat()}\n\n")
|
||||
f.write(f"**Repository:** {stats.get('repository_path', 'Unknown')}\n\n")
|
||||
|
||||
|
||||
# Summary statistics
|
||||
f.write("## Summary Statistics\n\n")
|
||||
f.write(f"- **Total Files:** {stats.get('total_files', 0)}\n")
|
||||
f.write(f"- **Total Contributors:** {stats.get('total_contributors', 0)}\n\n")
|
||||
|
||||
|
||||
# Categories
|
||||
if 'categories' in stats:
|
||||
if "categories" in stats:
|
||||
f.write("### Files by Category\n\n")
|
||||
for category, count in sorted(stats['categories'].items()):
|
||||
for category, count in sorted(stats["categories"].items()):
|
||||
f.write(f"- **{category}:** {count} files\n")
|
||||
f.write("\n")
|
||||
|
||||
|
||||
# Top contributors
|
||||
if 'contributor_file_counts' in stats:
|
||||
if "contributor_file_counts" in stats:
|
||||
f.write("### Top Contributors by Files Modified\n\n")
|
||||
sorted_contributors = sorted(
|
||||
stats['contributor_file_counts'].items(),
|
||||
key=lambda x: x[1],
|
||||
reverse=True
|
||||
stats["contributor_file_counts"].items(), key=lambda x: x[1], reverse=True
|
||||
)[:10]
|
||||
for contributor, count in sorted_contributors:
|
||||
f.write(f"- **{contributor}:** {count} files\n")
|
||||
f.write("\n")
|
||||
|
||||
|
||||
# Detailed file listing
|
||||
f.write("## Detailed File Inventory\n\n")
|
||||
f.write("| File Path | Category | Size (bytes) | Primary Author | Contributors |\n")
|
||||
f.write("|-----------|----------|--------------|----------------|-------------|\n")
|
||||
|
||||
for entry in sorted(self.inventory, key=lambda x: x['file_path']):
|
||||
contributors_str = ', '.join(entry['contributors'].keys())[:50]
|
||||
|
||||
for entry in sorted(self.inventory, key=lambda x: x["file_path"]):
|
||||
contributors_str = ", ".join(entry["contributors"].keys())[:50]
|
||||
if len(contributors_str) == 50:
|
||||
contributors_str += "..."
|
||||
|
||||
f.write(f"| {entry['file_path']} | {entry['category']} | "
|
||||
f"{entry['file_size_bytes']} | {entry['primary_author']} | "
|
||||
f"{contributors_str} |\n")
|
||||
|
||||
f.write(
|
||||
f"| {entry['file_path']} | {entry['category']} | "
|
||||
f"{entry['file_size_bytes']} | {entry['primary_author']} | "
|
||||
f"{contributors_str} |\n"
|
||||
)
|
||||
|
||||
|
||||
def main():
|
||||
parser = argparse.ArgumentParser(
|
||||
description="Generate legal file inventory for Basic Memory repository"
|
||||
)
|
||||
parser.add_argument(
|
||||
'--output', '-o',
|
||||
default='basic_memory_legal_inventory.csv',
|
||||
help='Output file path (default: basic_memory_legal_inventory.csv)'
|
||||
"--output",
|
||||
"-o",
|
||||
default="basic_memory_legal_inventory.csv",
|
||||
help="Output file path (default: basic_memory_legal_inventory.csv)",
|
||||
)
|
||||
parser.add_argument(
|
||||
'--format', '-f',
|
||||
choices=['csv', 'json', 'markdown'],
|
||||
default='csv',
|
||||
help='Output format (default: csv)'
|
||||
"--format",
|
||||
"-f",
|
||||
choices=["csv", "json", "markdown"],
|
||||
default="csv",
|
||||
help="Output format (default: csv)",
|
||||
)
|
||||
parser.add_argument(
|
||||
'--repo-path', '-r',
|
||||
default='.',
|
||||
help='Path to repository (default: current directory)'
|
||||
"--repo-path", "-r", default=".", help="Path to repository (default: current directory)"
|
||||
)
|
||||
|
||||
|
||||
args = parser.parse_args()
|
||||
|
||||
|
||||
# Initialize and run the inventory generator
|
||||
generator = FileInventoryGenerator(args.repo_path)
|
||||
generator.scan_repository()
|
||||
|
||||
|
||||
# Export in requested format
|
||||
if args.format == 'csv':
|
||||
if args.format == "csv":
|
||||
generator.export_csv(args.output)
|
||||
elif args.format == 'json':
|
||||
elif args.format == "json":
|
||||
generator.export_json(args.output)
|
||||
elif args.format == 'markdown':
|
||||
elif args.format == "markdown":
|
||||
generator.export_markdown(args.output)
|
||||
|
||||
|
||||
# Print summary
|
||||
stats = generator.get_summary_statistics()
|
||||
print("\n=== Legal File Inventory Complete ===")
|
||||
@@ -422,17 +471,16 @@ def main():
|
||||
print(f"Total files inventoried: {stats.get('total_files', 0)}")
|
||||
print(f"Total contributors identified: {stats.get('total_contributors', 0)}")
|
||||
print(f"Output written to: {args.output}")
|
||||
|
||||
|
||||
# Show top contributors
|
||||
if 'contributor_file_counts' in stats:
|
||||
if "contributor_file_counts" in stats:
|
||||
print("\nTop 5 contributors by files modified:")
|
||||
sorted_contributors = sorted(
|
||||
stats['contributor_file_counts'].items(),
|
||||
key=lambda x: x[1],
|
||||
reverse=True
|
||||
stats["contributor_file_counts"].items(), key=lambda x: x[1], reverse=True
|
||||
)[:5]
|
||||
for i, (contributor, count) in enumerate(sorted_contributors, 1):
|
||||
print(f" {i}. {contributor}: {count} files")
|
||||
|
||||
if __name__ == '__main__':
|
||||
main()
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
|
||||
+272
-227
@@ -27,215 +27,251 @@ from typing import Dict, List
|
||||
|
||||
class LegalInventoryGenerator:
|
||||
"""Generate comprehensive file inventory for legal documentation."""
|
||||
|
||||
|
||||
# Files to exclude from legal inventory
|
||||
EXCLUDED_PATTERNS = {
|
||||
# Generated/compiled files
|
||||
'*.pyc', '*.pyo', '*.pyd', '__pycache__',
|
||||
'*.so', '*.dylib', '*.dll',
|
||||
|
||||
"*.pyc",
|
||||
"*.pyo",
|
||||
"*.pyd",
|
||||
"__pycache__",
|
||||
"*.so",
|
||||
"*.dylib",
|
||||
"*.dll",
|
||||
# Build/cache directories
|
||||
'build/', 'dist/', '.eggs/', '*.egg-info/',
|
||||
'.coverage', '.pytest_cache/', '.mypy_cache/',
|
||||
'.ruff_cache/', '.tox/', 'venv/', '.venv/', 'env/', '.env/',
|
||||
'node_modules/', '.npm/', '.yarn/',
|
||||
|
||||
"build/",
|
||||
"dist/",
|
||||
".eggs/",
|
||||
"*.egg-info/",
|
||||
".coverage",
|
||||
".pytest_cache/",
|
||||
".mypy_cache/",
|
||||
".ruff_cache/",
|
||||
".tox/",
|
||||
"venv/",
|
||||
".venv/",
|
||||
"env/",
|
||||
".env/",
|
||||
"node_modules/",
|
||||
".npm/",
|
||||
".yarn/",
|
||||
# IDE and editor files
|
||||
'.vscode/', '.idea/', '*.swp', '*.swo', '*~',
|
||||
'.DS_Store', 'Thumbs.db',
|
||||
|
||||
".vscode/",
|
||||
".idea/",
|
||||
"*.swp",
|
||||
"*.swo",
|
||||
"*~",
|
||||
".DS_Store",
|
||||
"Thumbs.db",
|
||||
# Version control
|
||||
'.git/', '.gitignore',
|
||||
|
||||
".git/",
|
||||
".gitignore",
|
||||
# OS generated
|
||||
'desktop.ini', '*.tmp', '*.temp'
|
||||
"desktop.ini",
|
||||
"*.tmp",
|
||||
"*.temp",
|
||||
}
|
||||
|
||||
|
||||
# File categories for legal classification
|
||||
FILE_CATEGORIES = {
|
||||
'source_code': ['.py', '.pyx', '.pyi'],
|
||||
'documentation': ['.md', '.rst', '.txt'],
|
||||
'configuration': ['.toml', '.yaml', '.yml', '.json', '.ini', '.cfg'],
|
||||
'legal': ['LICENSE', 'COPYING', 'COPYRIGHT', '.md'],
|
||||
'build_deployment': ['Dockerfile', 'Makefile', 'justfile', '.sh'],
|
||||
'database': ['.sql', '.sqlite', '.db'],
|
||||
'templates': ['.j2', '.jinja2', '.hbs', '.handlebars'],
|
||||
'data': ['.csv', '.json', '.xml'],
|
||||
'other': [] # Catch-all for uncategorized files
|
||||
"source_code": [".py", ".pyx", ".pyi"],
|
||||
"documentation": [".md", ".rst", ".txt"],
|
||||
"configuration": [".toml", ".yaml", ".yml", ".json", ".ini", ".cfg"],
|
||||
"legal": ["LICENSE", "COPYING", "COPYRIGHT", ".md"],
|
||||
"build_deployment": ["Dockerfile", "Makefile", "justfile", ".sh"],
|
||||
"database": [".sql", ".sqlite", ".db"],
|
||||
"templates": [".j2", ".jinja2", ".hbs", ".handlebars"],
|
||||
"data": [".csv", ".json", ".xml"],
|
||||
"other": [], # Catch-all for uncategorized files
|
||||
}
|
||||
|
||||
|
||||
def __init__(self, repo_path: str = "."):
|
||||
"""Initialize generator with repository path."""
|
||||
self.repo_path = Path(repo_path).resolve()
|
||||
self.file_inventory: List[Dict] = []
|
||||
self.contributors: Dict[str, Dict] = defaultdict(lambda: {
|
||||
'email': '', 'commits': 0, 'lines_added': 0, 'files': set()
|
||||
})
|
||||
|
||||
self.contributors: Dict[str, Dict] = defaultdict(
|
||||
lambda: {"email": "", "commits": 0, "lines_added": 0, "files": set()}
|
||||
)
|
||||
|
||||
def should_exclude_file(self, file_path: Path) -> bool:
|
||||
"""Check if file should be excluded from inventory."""
|
||||
# Check if file is tracked by git (more efficient than check-ignore)
|
||||
try:
|
||||
rel_path = str(file_path.relative_to(self.repo_path))
|
||||
result = subprocess.run([
|
||||
'git', 'ls-files', '--error-unmatch', rel_path
|
||||
], capture_output=True, cwd=self.repo_path)
|
||||
|
||||
result = subprocess.run(
|
||||
["git", "ls-files", "--error-unmatch", rel_path],
|
||||
capture_output=True,
|
||||
cwd=self.repo_path,
|
||||
)
|
||||
|
||||
# If git ls-files returns non-zero, file is not tracked (likely ignored)
|
||||
if result.returncode != 0:
|
||||
return True
|
||||
|
||||
|
||||
except Exception:
|
||||
# Fallback to manual exclusion patterns if git fails
|
||||
pass
|
||||
|
||||
|
||||
# Additional manual exclusions for safety
|
||||
file_str = str(file_path.relative_to(self.repo_path))
|
||||
|
||||
|
||||
for pattern in self.EXCLUDED_PATTERNS:
|
||||
if pattern.endswith('/'):
|
||||
if pattern.endswith("/"):
|
||||
if any(part == pattern[:-1] for part in file_path.parts):
|
||||
return True
|
||||
elif '*' in pattern:
|
||||
elif "*" in pattern:
|
||||
import fnmatch
|
||||
|
||||
if fnmatch.fnmatch(file_str, pattern):
|
||||
return True
|
||||
else:
|
||||
if file_path.name == pattern or file_str == pattern:
|
||||
return True
|
||||
return False
|
||||
|
||||
|
||||
def categorize_file(self, file_path: Path) -> str:
|
||||
"""Categorize file based on extension and name."""
|
||||
suffix = file_path.suffix.lower()
|
||||
name = file_path.name.upper()
|
||||
|
||||
|
||||
# Check legal files by name first
|
||||
if any(legal in name for legal in ['LICENSE', 'COPYING', 'COPYRIGHT', 'CLA']):
|
||||
return 'legal'
|
||||
|
||||
if any(legal in name for legal in ["LICENSE", "COPYING", "COPYRIGHT", "CLA"]):
|
||||
return "legal"
|
||||
|
||||
# Check by extension
|
||||
for category, extensions in self.FILE_CATEGORIES.items():
|
||||
if suffix in extensions:
|
||||
return category
|
||||
|
||||
return 'other'
|
||||
|
||||
|
||||
return "other"
|
||||
|
||||
def get_file_hash(self, file_path: Path) -> str:
|
||||
"""Generate SHA-256 hash of file content."""
|
||||
try:
|
||||
with open(file_path, 'rb') as f:
|
||||
with open(file_path, "rb") as f:
|
||||
return hashlib.sha256(f.read()).hexdigest()
|
||||
except (IOError, OSError):
|
||||
return "ERROR_READING_FILE"
|
||||
|
||||
|
||||
def get_git_contributors(self, file_path: Path) -> List[Dict]:
|
||||
"""Get contributor information for a specific file."""
|
||||
try:
|
||||
rel_path = file_path.relative_to(self.repo_path)
|
||||
|
||||
|
||||
# Get contributors with line counts
|
||||
result = subprocess.run([
|
||||
'git', 'log', '--follow', '--pretty=format:%an|%ae|%ad|%H',
|
||||
'--date=short', '--', str(rel_path)
|
||||
], capture_output=True, text=True, cwd=self.repo_path)
|
||||
|
||||
result = subprocess.run(
|
||||
[
|
||||
"git",
|
||||
"log",
|
||||
"--follow",
|
||||
"--pretty=format:%an|%ae|%ad|%H",
|
||||
"--date=short",
|
||||
"--",
|
||||
str(rel_path),
|
||||
],
|
||||
capture_output=True,
|
||||
text=True,
|
||||
cwd=self.repo_path,
|
||||
)
|
||||
|
||||
if result.returncode != 0:
|
||||
return []
|
||||
|
||||
|
||||
contributors = []
|
||||
seen = set()
|
||||
|
||||
for line in result.stdout.strip().split('\n'):
|
||||
|
||||
for line in result.stdout.strip().split("\n"):
|
||||
if not line:
|
||||
continue
|
||||
|
||||
parts = line.split('|')
|
||||
|
||||
parts = line.split("|")
|
||||
if len(parts) >= 4:
|
||||
name, email, date, commit_hash = parts[:4]
|
||||
|
||||
|
||||
# Normalize author names/emails
|
||||
normalized_name = self.normalize_author_name(name, email)
|
||||
|
||||
|
||||
if normalized_name not in seen:
|
||||
contributors.append({
|
||||
'name': normalized_name,
|
||||
'email': email,
|
||||
'first_contribution': date,
|
||||
'commit_hash': commit_hash
|
||||
})
|
||||
contributors.append(
|
||||
{
|
||||
"name": normalized_name,
|
||||
"email": email,
|
||||
"first_contribution": date,
|
||||
"commit_hash": commit_hash,
|
||||
}
|
||||
)
|
||||
seen.add(normalized_name)
|
||||
|
||||
|
||||
return contributors
|
||||
|
||||
|
||||
except Exception as e:
|
||||
print(f"Warning: Could not get git info for {file_path}: {e}")
|
||||
return []
|
||||
|
||||
|
||||
def normalize_author_name(self, name: str, email: str) -> str:
|
||||
"""Normalize author names to handle multiple emails for same person."""
|
||||
# Known mappings for Basic Memory team
|
||||
name_mappings = {
|
||||
'phernandez': 'Paul Hernandez',
|
||||
'Paul Hernandez': 'Paul Hernandez',
|
||||
'drew-cain': 'Drew Cain',
|
||||
'Drew Cain': 'Drew Cain'
|
||||
"phernandez": "Paul Hernandez",
|
||||
"Paul Hernandez": "Paul Hernandez",
|
||||
"drew-cain": "Drew Cain",
|
||||
"Drew Cain": "Drew Cain",
|
||||
}
|
||||
|
||||
|
||||
# Handle GitHub bot accounts
|
||||
if 'bot' in name.lower() or 'claude' in name.lower():
|
||||
if "bot" in name.lower() or "claude" in name.lower():
|
||||
return f"{name} (AI Assistant)"
|
||||
|
||||
|
||||
return name_mappings.get(name, name)
|
||||
|
||||
|
||||
def get_file_stats(self, file_path: Path) -> Dict:
|
||||
"""Get comprehensive file statistics."""
|
||||
try:
|
||||
stat = file_path.stat()
|
||||
rel_path = file_path.relative_to(self.repo_path)
|
||||
|
||||
|
||||
# Basic file info
|
||||
file_info = {
|
||||
'path': str(rel_path),
|
||||
'name': file_path.name,
|
||||
'size_bytes': stat.st_size,
|
||||
'modified_time': datetime.fromtimestamp(stat.st_mtime).isoformat(),
|
||||
'category': self.categorize_file(file_path),
|
||||
'sha256_hash': self.get_file_hash(file_path)
|
||||
"path": str(rel_path),
|
||||
"name": file_path.name,
|
||||
"size_bytes": stat.st_size,
|
||||
"modified_time": datetime.fromtimestamp(stat.st_mtime).isoformat(),
|
||||
"category": self.categorize_file(file_path),
|
||||
"sha256_hash": self.get_file_hash(file_path),
|
||||
}
|
||||
|
||||
|
||||
# Git information
|
||||
contributors = self.get_git_contributors(file_path)
|
||||
file_info['contributors'] = contributors
|
||||
file_info['primary_author'] = contributors[0]['name'] if contributors else 'Unknown'
|
||||
file_info['contributor_count'] = len(contributors)
|
||||
|
||||
file_info["contributors"] = contributors
|
||||
file_info["primary_author"] = contributors[0]["name"] if contributors else "Unknown"
|
||||
file_info["contributor_count"] = len(contributors)
|
||||
|
||||
# Update global contributor stats
|
||||
for contrib in contributors:
|
||||
name = contrib['name']
|
||||
self.contributors[name]['email'] = contrib['email']
|
||||
self.contributors[name]['files'].add(str(rel_path))
|
||||
|
||||
name = contrib["name"]
|
||||
self.contributors[name]["email"] = contrib["email"]
|
||||
self.contributors[name]["files"].add(str(rel_path))
|
||||
|
||||
return file_info
|
||||
|
||||
|
||||
except Exception as e:
|
||||
print(f"Error processing {file_path}: {e}")
|
||||
return None
|
||||
|
||||
|
||||
def scan_repository(self) -> None:
|
||||
"""Scan repository and build file inventory."""
|
||||
print(f"Scanning repository: {self.repo_path}")
|
||||
|
||||
|
||||
# Get all git-tracked files first (much more efficient)
|
||||
try:
|
||||
result = subprocess.run([
|
||||
'git', 'ls-files'
|
||||
], capture_output=True, text=True, cwd=self.repo_path)
|
||||
|
||||
result = subprocess.run(
|
||||
["git", "ls-files"], capture_output=True, text=True, cwd=self.repo_path
|
||||
)
|
||||
|
||||
if result.returncode == 0:
|
||||
tracked_files = [self.repo_path / f for f in result.stdout.strip().split('\n') if f]
|
||||
tracked_files = [self.repo_path / f for f in result.stdout.strip().split("\n") if f]
|
||||
print(f"Found {len(tracked_files)} git-tracked files")
|
||||
|
||||
|
||||
for file_path in tracked_files:
|
||||
if file_path.is_file():
|
||||
file_info = self.get_file_stats(file_path)
|
||||
@@ -244,141 +280,150 @@ class LegalInventoryGenerator:
|
||||
else:
|
||||
print("Warning: Could not get git tracked files, falling back to directory scan")
|
||||
self._fallback_scan()
|
||||
|
||||
|
||||
except Exception as e:
|
||||
print(f"Warning: Git command failed ({e}), falling back to directory scan")
|
||||
self._fallback_scan()
|
||||
|
||||
|
||||
# Get global git stats
|
||||
self._get_global_git_stats()
|
||||
|
||||
|
||||
print(f"Processed {len(self.file_inventory)} files")
|
||||
print(f"Found {len(self.contributors)} contributors")
|
||||
|
||||
|
||||
def _fallback_scan(self) -> None:
|
||||
"""Fallback directory scan if git commands fail."""
|
||||
for file_path in self.repo_path.rglob('*'):
|
||||
for file_path in self.repo_path.rglob("*"):
|
||||
if file_path.is_file() and not self.should_exclude_file(file_path):
|
||||
file_info = self.get_file_stats(file_path)
|
||||
if file_info:
|
||||
self.file_inventory.append(file_info)
|
||||
|
||||
|
||||
def _get_global_git_stats(self) -> None:
|
||||
"""Get global contributor statistics from git."""
|
||||
try:
|
||||
# Get commit counts per author
|
||||
result = subprocess.run([
|
||||
'git', 'shortlog', '-sn', '--all'
|
||||
], capture_output=True, text=True, cwd=self.repo_path)
|
||||
|
||||
result = subprocess.run(
|
||||
["git", "shortlog", "-sn", "--all"],
|
||||
capture_output=True,
|
||||
text=True,
|
||||
cwd=self.repo_path,
|
||||
)
|
||||
|
||||
if result.returncode == 0:
|
||||
for line in result.stdout.strip().split('\n'):
|
||||
for line in result.stdout.strip().split("\n"):
|
||||
if line.strip():
|
||||
parts = line.strip().split('\t', 1)
|
||||
parts = line.strip().split("\t", 1)
|
||||
if len(parts) == 2:
|
||||
count, name = parts
|
||||
normalized_name = self.normalize_author_name(name, '')
|
||||
self.contributors[normalized_name]['commits'] = int(count)
|
||||
|
||||
normalized_name = self.normalize_author_name(name, "")
|
||||
self.contributors[normalized_name]["commits"] = int(count)
|
||||
|
||||
except Exception as e:
|
||||
print(f"Warning: Could not get global git stats: {e}")
|
||||
|
||||
|
||||
def generate_summary(self) -> Dict:
|
||||
"""Generate inventory summary statistics."""
|
||||
total_files = len(self.file_inventory)
|
||||
total_size = sum(f['size_bytes'] for f in self.file_inventory)
|
||||
|
||||
total_size = sum(f["size_bytes"] for f in self.file_inventory)
|
||||
|
||||
# Category breakdown
|
||||
categories = defaultdict(int)
|
||||
for file_info in self.file_inventory:
|
||||
categories[file_info['category']] += 1
|
||||
|
||||
categories[file_info["category"]] += 1
|
||||
|
||||
# Top contributors
|
||||
top_contributors = sorted(
|
||||
self.contributors.items(),
|
||||
key=lambda x: len(x[1]['files']),
|
||||
reverse=True
|
||||
self.contributors.items(), key=lambda x: len(x[1]["files"]), reverse=True
|
||||
)[:10]
|
||||
|
||||
|
||||
return {
|
||||
'scan_date': datetime.now().isoformat(),
|
||||
'repository_path': str(self.repo_path),
|
||||
'total_files': total_files,
|
||||
'total_size_bytes': total_size,
|
||||
'categories': dict(categories),
|
||||
'contributor_count': len(self.contributors),
|
||||
'top_contributors': [
|
||||
"scan_date": datetime.now().isoformat(),
|
||||
"repository_path": str(self.repo_path),
|
||||
"total_files": total_files,
|
||||
"total_size_bytes": total_size,
|
||||
"categories": dict(categories),
|
||||
"contributor_count": len(self.contributors),
|
||||
"top_contributors": [
|
||||
{
|
||||
'name': name,
|
||||
'file_count': len(stats['files']),
|
||||
'commit_count': stats['commits'],
|
||||
'email': stats['email']
|
||||
"name": name,
|
||||
"file_count": len(stats["files"]),
|
||||
"commit_count": stats["commits"],
|
||||
"email": stats["email"],
|
||||
}
|
||||
for name, stats in top_contributors
|
||||
]
|
||||
],
|
||||
}
|
||||
|
||||
|
||||
def export_csv(self, output_path: str) -> None:
|
||||
"""Export inventory to CSV format."""
|
||||
with open(output_path, 'w', newline='', encoding='utf-8') as csvfile:
|
||||
with open(output_path, "w", newline="", encoding="utf-8") as csvfile:
|
||||
fieldnames = [
|
||||
'path', 'name', 'category', 'size_bytes', 'modified_time',
|
||||
'primary_author', 'contributor_count', 'contributors_list',
|
||||
'sha256_hash'
|
||||
"path",
|
||||
"name",
|
||||
"category",
|
||||
"size_bytes",
|
||||
"modified_time",
|
||||
"primary_author",
|
||||
"contributor_count",
|
||||
"contributors_list",
|
||||
"sha256_hash",
|
||||
]
|
||||
|
||||
|
||||
writer = csv.DictWriter(csvfile, fieldnames=fieldnames)
|
||||
writer.writeheader()
|
||||
|
||||
for file_info in sorted(self.file_inventory, key=lambda x: x['path']):
|
||||
contributors_list = '; '.join([
|
||||
f"{c['name']} ({c['email']})" for c in file_info['contributors']
|
||||
])
|
||||
|
||||
writer.writerow({
|
||||
'path': file_info['path'],
|
||||
'name': file_info['name'],
|
||||
'category': file_info['category'],
|
||||
'size_bytes': file_info['size_bytes'],
|
||||
'modified_time': file_info['modified_time'],
|
||||
'primary_author': file_info['primary_author'],
|
||||
'contributor_count': file_info['contributor_count'],
|
||||
'contributors_list': contributors_list,
|
||||
'sha256_hash': file_info['sha256_hash']
|
||||
})
|
||||
|
||||
|
||||
for file_info in sorted(self.file_inventory, key=lambda x: x["path"]):
|
||||
contributors_list = "; ".join(
|
||||
[f"{c['name']} ({c['email']})" for c in file_info["contributors"]]
|
||||
)
|
||||
|
||||
writer.writerow(
|
||||
{
|
||||
"path": file_info["path"],
|
||||
"name": file_info["name"],
|
||||
"category": file_info["category"],
|
||||
"size_bytes": file_info["size_bytes"],
|
||||
"modified_time": file_info["modified_time"],
|
||||
"primary_author": file_info["primary_author"],
|
||||
"contributor_count": file_info["contributor_count"],
|
||||
"contributors_list": contributors_list,
|
||||
"sha256_hash": file_info["sha256_hash"],
|
||||
}
|
||||
)
|
||||
|
||||
print(f"CSV export saved to: {output_path}")
|
||||
|
||||
|
||||
def export_json(self, output_path: str) -> None:
|
||||
"""Export inventory to JSON format."""
|
||||
# Convert sets to lists for JSON serialization
|
||||
contributors_serializable = {}
|
||||
for name, stats in self.contributors.items():
|
||||
contributors_serializable[name] = {
|
||||
'email': stats['email'],
|
||||
'commits': stats['commits'],
|
||||
'lines_added': stats['lines_added'],
|
||||
'files': list(stats['files'])
|
||||
"email": stats["email"],
|
||||
"commits": stats["commits"],
|
||||
"lines_added": stats["lines_added"],
|
||||
"files": list(stats["files"]),
|
||||
}
|
||||
|
||||
|
||||
data = {
|
||||
'summary': self.generate_summary(),
|
||||
'files': self.file_inventory,
|
||||
'contributors': contributors_serializable
|
||||
"summary": self.generate_summary(),
|
||||
"files": self.file_inventory,
|
||||
"contributors": contributors_serializable,
|
||||
}
|
||||
|
||||
with open(output_path, 'w', encoding='utf-8') as jsonfile:
|
||||
|
||||
with open(output_path, "w", encoding="utf-8") as jsonfile:
|
||||
json.dump(data, jsonfile, indent=2, ensure_ascii=False)
|
||||
|
||||
|
||||
print(f"JSON export saved to: {output_path}")
|
||||
|
||||
|
||||
def export_markdown(self, output_path: str) -> None:
|
||||
"""Export inventory to Markdown format for legal documentation."""
|
||||
summary = self.generate_summary()
|
||||
|
||||
with open(output_path, 'w', encoding='utf-8') as mdfile:
|
||||
|
||||
with open(output_path, "w", encoding="utf-8") as mdfile:
|
||||
mdfile.write("# Basic Memory - Legal File Inventory\n\n")
|
||||
|
||||
|
||||
# Summary section
|
||||
mdfile.write("## Summary\n\n")
|
||||
mdfile.write(f"**Scan Date:** {summary['scan_date']}\n")
|
||||
@@ -386,42 +431,42 @@ class LegalInventoryGenerator:
|
||||
mdfile.write(f"**Total Files:** {summary['total_files']:,}\n")
|
||||
mdfile.write(f"**Total Size:** {summary['total_size_bytes']:,} bytes\n")
|
||||
mdfile.write(f"**Contributors:** {summary['contributor_count']}\n\n")
|
||||
|
||||
|
||||
# Category breakdown
|
||||
mdfile.write("## File Categories\n\n")
|
||||
for category, count in sorted(summary['categories'].items()):
|
||||
for category, count in sorted(summary["categories"].items()):
|
||||
mdfile.write(f"- **{category.replace('_', ' ').title()}:** {count} files\n")
|
||||
mdfile.write("\n")
|
||||
|
||||
|
||||
# Top contributors
|
||||
mdfile.write("## Contributors\n\n")
|
||||
for contrib in summary['top_contributors']:
|
||||
for contrib in summary["top_contributors"]:
|
||||
mdfile.write(f"- **{contrib['name']}** ({contrib['email']}): ")
|
||||
mdfile.write(f"{contrib['file_count']} files, {contrib['commit_count']} commits\n")
|
||||
mdfile.write("\n")
|
||||
|
||||
|
||||
# Detailed file listing by category
|
||||
mdfile.write("## Detailed File Inventory\n\n")
|
||||
|
||||
for category in sorted(summary['categories'].keys()):
|
||||
category_files = [f for f in self.file_inventory if f['category'] == category]
|
||||
|
||||
for category in sorted(summary["categories"].keys()):
|
||||
category_files = [f for f in self.file_inventory if f["category"] == category]
|
||||
if not category_files:
|
||||
continue
|
||||
|
||||
|
||||
mdfile.write(f"### {category.replace('_', ' ').title()}\n\n")
|
||||
|
||||
for file_info in sorted(category_files, key=lambda x: x['path']):
|
||||
|
||||
for file_info in sorted(category_files, key=lambda x: x["path"]):
|
||||
mdfile.write(f"**{file_info['path']}**\n")
|
||||
mdfile.write(f"- Primary Author: {file_info['primary_author']}\n")
|
||||
mdfile.write(f"- Contributors: {file_info['contributor_count']}\n")
|
||||
mdfile.write(f"- Size: {file_info['size_bytes']:,} bytes\n")
|
||||
|
||||
if file_info['contributors']:
|
||||
contributors_str = ', '.join([c['name'] for c in file_info['contributors']])
|
||||
|
||||
if file_info["contributors"]:
|
||||
contributors_str = ", ".join([c["name"] for c in file_info["contributors"]])
|
||||
mdfile.write(f"- All Contributors: {contributors_str}\n")
|
||||
|
||||
|
||||
mdfile.write(f"- SHA-256: `{file_info['sha256_hash']}`\n\n")
|
||||
|
||||
|
||||
print(f"Markdown export saved to: {output_path}")
|
||||
|
||||
|
||||
@@ -431,53 +476,53 @@ def main():
|
||||
description="Generate legal file inventory for Basic Memory repository"
|
||||
)
|
||||
parser.add_argument(
|
||||
'--repo-path', '-r',
|
||||
default='.',
|
||||
help='Path to repository (default: current directory)'
|
||||
"--repo-path", "-r", default=".", help="Path to repository (default: current directory)"
|
||||
)
|
||||
parser.add_argument(
|
||||
'--output-dir', '-o',
|
||||
default='./legal_inventory',
|
||||
help='Output directory for reports (default: ./legal_inventory)'
|
||||
"--output-dir",
|
||||
"-o",
|
||||
default="./legal_inventory",
|
||||
help="Output directory for reports (default: ./legal_inventory)",
|
||||
)
|
||||
parser.add_argument(
|
||||
'--formats', '-f',
|
||||
nargs='+',
|
||||
choices=['csv', 'json', 'markdown', 'all'],
|
||||
default=['all'],
|
||||
help='Output formats to generate (default: all)'
|
||||
"--formats",
|
||||
"-f",
|
||||
nargs="+",
|
||||
choices=["csv", "json", "markdown", "all"],
|
||||
default=["all"],
|
||||
help="Output formats to generate (default: all)",
|
||||
)
|
||||
|
||||
|
||||
args = parser.parse_args()
|
||||
|
||||
|
||||
# Create output directory
|
||||
output_dir = Path(args.output_dir)
|
||||
output_dir.mkdir(exist_ok=True)
|
||||
|
||||
|
||||
# Generate inventory
|
||||
generator = LegalInventoryGenerator(args.repo_path)
|
||||
generator.scan_repository()
|
||||
|
||||
|
||||
# Determine formats to export
|
||||
formats = args.formats
|
||||
if 'all' in formats:
|
||||
formats = ['csv', 'json', 'markdown']
|
||||
|
||||
if "all" in formats:
|
||||
formats = ["csv", "json", "markdown"]
|
||||
|
||||
# Export in requested formats
|
||||
timestamp = datetime.now().strftime('%Y%m%d_%H%M%S')
|
||||
|
||||
if 'csv' in formats:
|
||||
generator.export_csv(output_dir / f'basic_memory_inventory_{timestamp}.csv')
|
||||
|
||||
if 'json' in formats:
|
||||
generator.export_json(output_dir / f'basic_memory_inventory_{timestamp}.json')
|
||||
|
||||
if 'markdown' in formats:
|
||||
generator.export_markdown(output_dir / f'basic_memory_inventory_{timestamp}.md')
|
||||
|
||||
timestamp = datetime.now().strftime("%Y%m%d_%H%M%S")
|
||||
|
||||
if "csv" in formats:
|
||||
generator.export_csv(output_dir / f"basic_memory_inventory_{timestamp}.csv")
|
||||
|
||||
if "json" in formats:
|
||||
generator.export_json(output_dir / f"basic_memory_inventory_{timestamp}.json")
|
||||
|
||||
if "markdown" in formats:
|
||||
generator.export_markdown(output_dir / f"basic_memory_inventory_{timestamp}.md")
|
||||
|
||||
print("\nLegal inventory generation complete!")
|
||||
print(f"Output saved to: {output_dir}")
|
||||
|
||||
|
||||
if __name__ == '__main__':
|
||||
main()
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
|
||||
@@ -20,14 +20,14 @@ depends_on: Union[str, Sequence[str], None] = None
|
||||
|
||||
def upgrade() -> None:
|
||||
"""Re-establish foreign key constraints that were lost during project table recreation.
|
||||
|
||||
|
||||
The migration 647e7a75e2cd recreated the project table but did not re-establish
|
||||
the foreign key constraint from entity.project_id to project.id, causing
|
||||
foreign key constraint failures when trying to delete projects with related entities.
|
||||
"""
|
||||
# SQLite doesn't allow adding foreign key constraints to existing tables easily
|
||||
# We need to be careful and handle the case where the constraint might already exist
|
||||
|
||||
|
||||
with op.batch_alter_table("entity", schema=None) as batch_op:
|
||||
# Try to drop existing foreign key constraint (may not exist)
|
||||
try:
|
||||
@@ -35,19 +35,15 @@ def upgrade() -> None:
|
||||
except Exception:
|
||||
# Constraint may not exist, which is fine - we'll create it next
|
||||
pass
|
||||
|
||||
|
||||
# Add the foreign key constraint with CASCADE DELETE
|
||||
# This ensures that when a project is deleted, all related entities are also deleted
|
||||
batch_op.create_foreign_key(
|
||||
"fk_entity_project_id",
|
||||
"project",
|
||||
["project_id"],
|
||||
["id"],
|
||||
ondelete="CASCADE"
|
||||
"fk_entity_project_id", "project", ["project_id"], ["id"], ondelete="CASCADE"
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
"""Remove the foreign key constraint."""
|
||||
with op.batch_alter_table("entity", schema=None) as batch_op:
|
||||
batch_op.drop_constraint("fk_entity_project_id", type_="foreignkey")
|
||||
batch_op.drop_constraint("fk_entity_project_id", type_="foreignkey")
|
||||
|
||||
@@ -240,40 +240,37 @@ async def update_frontmatter(path: FilePath, updates: Dict[str, Any]) -> str:
|
||||
def dump_frontmatter(post: frontmatter.Post) -> str:
|
||||
"""
|
||||
Serialize frontmatter.Post to markdown with Obsidian-compatible YAML format.
|
||||
|
||||
|
||||
This function ensures that tags are formatted as YAML lists instead of JSON arrays:
|
||||
|
||||
|
||||
Good (Obsidian compatible):
|
||||
---
|
||||
tags:
|
||||
- system
|
||||
- overview
|
||||
- overview
|
||||
- reference
|
||||
---
|
||||
|
||||
|
||||
Bad (current behavior):
|
||||
---
|
||||
tags: ["system", "overview", "reference"]
|
||||
---
|
||||
|
||||
|
||||
Args:
|
||||
post: frontmatter.Post object to serialize
|
||||
|
||||
|
||||
Returns:
|
||||
String containing markdown with properly formatted YAML frontmatter
|
||||
"""
|
||||
"""
|
||||
if not post.metadata:
|
||||
# No frontmatter, just return content
|
||||
return post.content
|
||||
|
||||
|
||||
# Serialize YAML with block style for lists
|
||||
yaml_str = yaml.dump(
|
||||
post.metadata,
|
||||
sort_keys=False,
|
||||
allow_unicode=True,
|
||||
default_flow_style=False
|
||||
post.metadata, sort_keys=False, allow_unicode=True, default_flow_style=False
|
||||
)
|
||||
|
||||
|
||||
# Construct the final markdown with frontmatter
|
||||
if post.content:
|
||||
return f"---\n{yaml_str}---\n\n{post.content}"
|
||||
@@ -297,4 +294,3 @@ def sanitize_for_filename(text: str, replacement: str = "-") -> str:
|
||||
text = re.sub(f"{re.escape(replacement)}+", replacement, text)
|
||||
|
||||
return text.strip(replacement)
|
||||
|
||||
|
||||
@@ -9,6 +9,7 @@ from markdown_it.token import Token
|
||||
def is_observation(token: Token) -> bool:
|
||||
"""Check if token looks like our observation format."""
|
||||
import re
|
||||
|
||||
if token.type != "inline": # pragma: no cover
|
||||
return False
|
||||
# Use token.tag which contains the actual content for test tokens, fallback to content
|
||||
@@ -18,15 +19,15 @@ def is_observation(token: Token) -> bool:
|
||||
# if it's a markdown_task, return false
|
||||
if content.startswith("[ ]") or content.startswith("[x]") or content.startswith("[-]"):
|
||||
return False
|
||||
|
||||
|
||||
# Exclude markdown links: [text](url)
|
||||
if re.match(r"^\[.*?\]\(.*?\)$", content):
|
||||
return False
|
||||
|
||||
|
||||
# Exclude wiki links: [[text]]
|
||||
if re.match(r"^\[\[.*?\]\]$", content):
|
||||
return False
|
||||
|
||||
|
||||
# Check for proper observation format: [category] content
|
||||
match = re.match(r"^\[([^\[\]()]+)\]\s+(.+)", content)
|
||||
has_tags = "#" in content
|
||||
@@ -36,9 +37,10 @@ def is_observation(token: Token) -> bool:
|
||||
def parse_observation(token: Token) -> Dict[str, Any]:
|
||||
"""Extract observation parts from token."""
|
||||
import re
|
||||
|
||||
# Use token.tag which contains the actual content for test tokens, fallback to content
|
||||
content = (token.tag or token.content).strip()
|
||||
|
||||
|
||||
# Parse [category] with regex
|
||||
match = re.match(r"^\[([^\[\]()]+)\]\s+(.+)", content)
|
||||
category = None
|
||||
@@ -50,7 +52,7 @@ def parse_observation(token: Token) -> Dict[str, Any]:
|
||||
empty_match = re.match(r"^\[\]\s+(.+)", content)
|
||||
if empty_match:
|
||||
content = empty_match.group(1).strip()
|
||||
|
||||
|
||||
# Parse (context)
|
||||
context = None
|
||||
if content.endswith(")"):
|
||||
@@ -58,7 +60,7 @@ def parse_observation(token: Token) -> Dict[str, Any]:
|
||||
if start != -1:
|
||||
context = content[start + 1 : -1].strip()
|
||||
content = content[:start].strip()
|
||||
|
||||
|
||||
# Extract tags and keep original content
|
||||
tags = []
|
||||
parts = content.split()
|
||||
@@ -69,7 +71,7 @@ def parse_observation(token: Token) -> Dict[str, Any]:
|
||||
tags.extend(subtags)
|
||||
else:
|
||||
tags.append(part[1:])
|
||||
|
||||
|
||||
return {
|
||||
"category": category,
|
||||
"content": content,
|
||||
|
||||
@@ -17,6 +17,7 @@ from basic_memory.schemas.memory import (
|
||||
|
||||
type StringOrInt = str | int
|
||||
|
||||
|
||||
@mcp.tool(
|
||||
description="""Build context from a memory:// URI to continue conversations naturally.
|
||||
|
||||
@@ -81,15 +82,16 @@ async def build_context(
|
||||
build_context("memory://specs/search", project="work-project")
|
||||
"""
|
||||
logger.info(f"Building context from {url}")
|
||||
|
||||
|
||||
# Convert string depth to integer if needed
|
||||
if isinstance(depth, str):
|
||||
try:
|
||||
depth = int(depth)
|
||||
except ValueError:
|
||||
from mcp.server.fastmcp.exceptions import ToolError
|
||||
|
||||
raise ToolError(f"Invalid depth parameter: '{depth}' is not a valid integer")
|
||||
|
||||
|
||||
# URL is already validated and normalized by MemoryUrl type annotation
|
||||
|
||||
# Get the active project first to check project-specific sync status
|
||||
|
||||
@@ -223,7 +223,8 @@ async def set_default_project(project_name: str, ctx: Context | None = None) ->
|
||||
|
||||
# Call API to set default project using URL encoding for special characters
|
||||
from urllib.parse import quote
|
||||
encoded_name = quote(project_name, safe='')
|
||||
|
||||
encoded_name = quote(project_name, safe="")
|
||||
response = await call_put(client, f"/projects/{encoded_name}/default")
|
||||
status_response = ProjectStatusResponse.model_validate(response.json())
|
||||
|
||||
@@ -337,7 +338,7 @@ async def delete_project(project_name: str, ctx: Context | None = None) -> str:
|
||||
if p.name.lower() == project_name.lower():
|
||||
target_project = p
|
||||
break
|
||||
|
||||
|
||||
if not target_project:
|
||||
available_projects = [p.name for p in project_list.projects]
|
||||
raise ValueError(
|
||||
@@ -346,7 +347,8 @@ async def delete_project(project_name: str, ctx: Context | None = None) -> str:
|
||||
|
||||
# Call API to delete project using URL encoding for special characters
|
||||
from urllib.parse import quote
|
||||
encoded_name = quote(target_project.name, safe='')
|
||||
|
||||
encoded_name = quote(target_project.name, safe="")
|
||||
response = await call_delete(client, f"/projects/{encoded_name}")
|
||||
status_response = ProjectStatusResponse.model_validate(response.json())
|
||||
|
||||
|
||||
@@ -60,8 +60,10 @@ async def read_note(
|
||||
# We need to check both the raw identifier and the processed path
|
||||
processed_path = memory_url_path(identifier)
|
||||
project_path = active_project.home
|
||||
|
||||
if not validate_project_path(identifier, project_path) or not validate_project_path(processed_path, project_path):
|
||||
|
||||
if not validate_project_path(identifier, project_path) or not validate_project_path(
|
||||
processed_path, project_path
|
||||
):
|
||||
logger.warning(
|
||||
"Attempted path traversal attack blocked",
|
||||
identifier=identifier,
|
||||
|
||||
@@ -74,8 +74,14 @@ class Entity(Base):
|
||||
checksum: Mapped[Optional[str]] = mapped_column(String, nullable=True)
|
||||
|
||||
# Metadata and tracking
|
||||
created_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), default=lambda: datetime.now().astimezone())
|
||||
updated_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), default=lambda: datetime.now().astimezone(), onupdate=lambda: datetime.now().astimezone())
|
||||
created_at: Mapped[datetime] = mapped_column(
|
||||
DateTime(timezone=True), default=lambda: datetime.now().astimezone()
|
||||
)
|
||||
updated_at: Mapped[datetime] = mapped_column(
|
||||
DateTime(timezone=True),
|
||||
default=lambda: datetime.now().astimezone(),
|
||||
onupdate=lambda: datetime.now().astimezone(),
|
||||
)
|
||||
|
||||
# Relationships
|
||||
project = relationship("Project", back_populates="entities")
|
||||
@@ -104,15 +110,15 @@ class Entity(Base):
|
||||
def is_markdown(self):
|
||||
"""Check if the entity is a markdown file."""
|
||||
return self.content_type == "text/markdown"
|
||||
|
||||
|
||||
def __getattribute__(self, name):
|
||||
"""Override attribute access to ensure datetime fields are timezone-aware."""
|
||||
value = super().__getattribute__(name)
|
||||
|
||||
|
||||
# Ensure datetime fields are timezone-aware
|
||||
if name in ('created_at', 'updated_at') and isinstance(value, datetime):
|
||||
if name in ("created_at", "updated_at") and isinstance(value, datetime):
|
||||
return ensure_timezone_aware(value)
|
||||
|
||||
|
||||
return value
|
||||
|
||||
def __repr__(self) -> str:
|
||||
|
||||
@@ -52,9 +52,13 @@ class Project(Base):
|
||||
is_default: Mapped[Optional[bool]] = mapped_column(Boolean, default=None, nullable=True)
|
||||
|
||||
# Timestamps
|
||||
created_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), default=lambda: datetime.now(UTC))
|
||||
created_at: Mapped[datetime] = mapped_column(
|
||||
DateTime(timezone=True), default=lambda: datetime.now(UTC)
|
||||
)
|
||||
updated_at: Mapped[datetime] = mapped_column(
|
||||
DateTime(timezone=True), default=lambda: datetime.now(UTC), onupdate=lambda: datetime.now(UTC)
|
||||
DateTime(timezone=True),
|
||||
default=lambda: datetime.now(UTC),
|
||||
onupdate=lambda: datetime.now(UTC),
|
||||
)
|
||||
|
||||
# Define relationships to entities, observations, and relations
|
||||
|
||||
@@ -62,7 +62,7 @@ class SearchIndexRow:
|
||||
|
||||
# Normalize path separators to handle both Windows (\) and Unix (/) paths
|
||||
normalized_path = Path(self.file_path).as_posix()
|
||||
|
||||
|
||||
# Split the path by slashes
|
||||
parts = normalized_path.split("/")
|
||||
|
||||
@@ -527,7 +527,9 @@ class SearchRepository:
|
||||
async with db.scoped_session(self.session_maker) as session:
|
||||
# Delete existing record if any
|
||||
await session.execute(
|
||||
text("DELETE FROM search_index WHERE permalink = :permalink AND project_id = :project_id"),
|
||||
text(
|
||||
"DELETE FROM search_index WHERE permalink = :permalink AND project_id = :project_id"
|
||||
),
|
||||
{"permalink": search_index_row.permalink, "project_id": self.project_id},
|
||||
)
|
||||
|
||||
|
||||
@@ -71,7 +71,7 @@ def parse_timeframe(timeframe: str) -> datetime:
|
||||
parsed = parse(timeframe)
|
||||
if not parsed:
|
||||
raise ValueError(f"Could not parse timeframe: {timeframe}")
|
||||
|
||||
|
||||
# If the parsed datetime is naive, make it timezone-aware in local system timezone
|
||||
if parsed.tzinfo is None:
|
||||
return parsed.astimezone()
|
||||
|
||||
@@ -117,7 +117,7 @@ def memory_url_path(url: memory_url) -> str: # pyright: ignore
|
||||
|
||||
class EntitySummary(BaseModel):
|
||||
"""Simplified entity representation."""
|
||||
|
||||
|
||||
model_config = ConfigDict(json_encoders={datetime: lambda dt: dt.isoformat()})
|
||||
|
||||
type: Literal["entity"] = "entity"
|
||||
@@ -130,7 +130,7 @@ class EntitySummary(BaseModel):
|
||||
|
||||
class RelationSummary(BaseModel):
|
||||
"""Simplified relation representation."""
|
||||
|
||||
|
||||
model_config = ConfigDict(json_encoders={datetime: lambda dt: dt.isoformat()})
|
||||
|
||||
type: Literal["relation"] = "relation"
|
||||
@@ -145,7 +145,7 @@ class RelationSummary(BaseModel):
|
||||
|
||||
class ObservationSummary(BaseModel):
|
||||
"""Simplified observation representation."""
|
||||
|
||||
|
||||
model_config = ConfigDict(json_encoders={datetime: lambda dt: dt.isoformat()})
|
||||
|
||||
type: Literal["observation"] = "observation"
|
||||
@@ -159,7 +159,7 @@ class ObservationSummary(BaseModel):
|
||||
|
||||
class MemoryMetadata(BaseModel):
|
||||
"""Simplified response metadata."""
|
||||
|
||||
|
||||
model_config = ConfigDict(json_encoders={datetime: lambda dt: dt.isoformat()})
|
||||
|
||||
uri: Optional[str] = None
|
||||
@@ -178,8 +178,8 @@ class ContextResult(BaseModel):
|
||||
"""Context result containing a primary item with its observations and related items."""
|
||||
|
||||
primary_result: Annotated[
|
||||
Union[EntitySummary, RelationSummary, ObservationSummary],
|
||||
Field(discriminator="type", description="Primary item")
|
||||
Union[EntitySummary, RelationSummary, ObservationSummary],
|
||||
Field(discriminator="type", description="Primary item"),
|
||||
]
|
||||
|
||||
observations: Sequence[ObservationSummary] = Field(
|
||||
@@ -188,8 +188,7 @@ class ContextResult(BaseModel):
|
||||
|
||||
related_results: Sequence[
|
||||
Annotated[
|
||||
Union[EntitySummary, RelationSummary, ObservationSummary],
|
||||
Field(discriminator="type")
|
||||
Union[EntitySummary, RelationSummary, ObservationSummary], Field(discriminator="type")
|
||||
]
|
||||
] = Field(description="Related items", default_factory=list)
|
||||
|
||||
|
||||
@@ -246,7 +246,11 @@ class ContextService:
|
||||
values = ", ".join([f"('{t}', {i})" for t, i in type_id_pairs])
|
||||
|
||||
# Parameters for bindings - include project_id for security filtering
|
||||
params = {"max_depth": max_depth, "max_results": max_results, "project_id": self.search_repository.project_id}
|
||||
params = {
|
||||
"max_depth": max_depth,
|
||||
"max_results": max_results,
|
||||
"project_id": self.search_repository.project_id,
|
||||
}
|
||||
|
||||
# Build date and timeframe filters conditionally based on since parameter
|
||||
if since:
|
||||
@@ -258,7 +262,7 @@ class ContextService:
|
||||
date_filter = ""
|
||||
relation_date_filter = ""
|
||||
timeframe_condition = ""
|
||||
|
||||
|
||||
# Add project filtering for security - ensure all entities and relations belong to the same project
|
||||
project_filter = "AND e.project_id = :project_id"
|
||||
relation_project_filter = "AND e_from.project_id = :project_id"
|
||||
|
||||
@@ -9,7 +9,12 @@ from loguru import logger
|
||||
from sqlalchemy.exc import IntegrityError
|
||||
|
||||
from basic_memory.config import ProjectConfig, BasicMemoryConfig
|
||||
from basic_memory.file_utils import has_frontmatter, parse_frontmatter, remove_frontmatter, dump_frontmatter
|
||||
from basic_memory.file_utils import (
|
||||
has_frontmatter,
|
||||
parse_frontmatter,
|
||||
remove_frontmatter,
|
||||
dump_frontmatter,
|
||||
)
|
||||
from basic_memory.markdown import EntityMarkdown
|
||||
from basic_memory.markdown.entity_parser import EntityParser
|
||||
from basic_memory.markdown.utils import entity_model_from_markdown, schema_to_markdown
|
||||
|
||||
@@ -288,9 +288,13 @@ class WatchService:
|
||||
full_path = directory / path
|
||||
if full_path.exists() and full_path.is_file():
|
||||
# File still exists despite DELETE event - treat as modification
|
||||
logger.debug("File exists despite DELETE event, treating as modification", path=path)
|
||||
logger.debug(
|
||||
"File exists despite DELETE event, treating as modification", path=path
|
||||
)
|
||||
entity, checksum = await sync_service.sync_file(path, new=False)
|
||||
self.state.add_event(path=path, action="modified", status="success", checksum=checksum)
|
||||
self.state.add_event(
|
||||
path=path, action="modified", status="success", checksum=checksum
|
||||
)
|
||||
self.console.print(f"[yellow]✎[/yellow] {path} (atomic write)")
|
||||
logger.info(f"atomic write detected: {path}")
|
||||
processed.add(path)
|
||||
@@ -302,10 +306,12 @@ class WatchService:
|
||||
entity = await sync_service.entity_repository.get_by_file_path(path)
|
||||
if entity is None:
|
||||
# No entity means this was likely a directory - skip it
|
||||
logger.debug(f"Skipping deleted path with no entity (likely directory), path={path}")
|
||||
logger.debug(
|
||||
f"Skipping deleted path with no entity (likely directory), path={path}"
|
||||
)
|
||||
processed.add(path)
|
||||
continue
|
||||
|
||||
|
||||
# File truly deleted
|
||||
logger.debug("Processing deleted file", path=path)
|
||||
await sync_service.handle_delete(path)
|
||||
|
||||
@@ -223,7 +223,8 @@ def parse_tags(tags: Union[List[str], str, None]) -> List[str]:
|
||||
if isinstance(tags, str):
|
||||
# Check if it's a JSON array string (common issue from AI assistants)
|
||||
import json
|
||||
if tags.strip().startswith('[') and tags.strip().endswith(']'):
|
||||
|
||||
if tags.strip().startswith("[") and tags.strip().endswith("]"):
|
||||
try:
|
||||
# Try to parse as JSON array
|
||||
parsed_json = json.loads(tags)
|
||||
@@ -233,7 +234,7 @@ def parse_tags(tags: Union[List[str], str, None]) -> List[str]:
|
||||
except json.JSONDecodeError:
|
||||
# Not valid JSON, fall through to comma-separated parsing
|
||||
pass
|
||||
|
||||
|
||||
# Split by comma, strip whitespace, then strip leading '#' characters
|
||||
return [tag.strip().lstrip("#") for tag in tags.split(",") if tag and tag.strip()]
|
||||
|
||||
|
||||
@@ -149,7 +149,8 @@ async def test_get_resource_observation(client, project_config, entity_repositor
|
||||
assert response.status_code == 200
|
||||
assert response.headers["content-type"] == "text/markdown; charset=utf-8"
|
||||
assert (
|
||||
normalize_newlines("""
|
||||
normalize_newlines(
|
||||
"""
|
||||
---
|
||||
title: Test Entity
|
||||
type: test
|
||||
@@ -159,7 +160,8 @@ permalink: test/test-entity
|
||||
# Test Content
|
||||
|
||||
- [note] an observation.
|
||||
""".strip())
|
||||
""".strip()
|
||||
)
|
||||
in response.text
|
||||
)
|
||||
|
||||
@@ -197,7 +199,8 @@ async def test_get_resource_entities(client, project_config, entity_repository,
|
||||
assert response.status_code == 200
|
||||
assert response.headers["content-type"] == "text/markdown; charset=utf-8"
|
||||
assert (
|
||||
normalize_newlines(f"""
|
||||
normalize_newlines(
|
||||
f"""
|
||||
--- memory://test/test-entity {entity1.updated_at.isoformat()} {entity1.checksum[:8]}
|
||||
|
||||
# Test Content
|
||||
@@ -207,7 +210,8 @@ async def test_get_resource_entities(client, project_config, entity_repository,
|
||||
# Related Content
|
||||
- links to [[Test Entity]]
|
||||
|
||||
""".strip())
|
||||
""".strip()
|
||||
)
|
||||
in response.text
|
||||
)
|
||||
|
||||
@@ -250,7 +254,8 @@ async def test_get_resource_entities_pagination(
|
||||
assert response.status_code == 200
|
||||
assert response.headers["content-type"] == "text/markdown; charset=utf-8"
|
||||
assert (
|
||||
normalize_newlines("""
|
||||
normalize_newlines(
|
||||
"""
|
||||
---
|
||||
title: Related Entity
|
||||
type: test
|
||||
@@ -259,7 +264,8 @@ permalink: test/related-entity
|
||||
|
||||
# Related Content
|
||||
- links to [[Test Entity]]
|
||||
""".strip())
|
||||
""".strip()
|
||||
)
|
||||
in response.text
|
||||
)
|
||||
|
||||
@@ -298,7 +304,8 @@ async def test_get_resource_relation(client, project_config, entity_repository,
|
||||
assert response.status_code == 200
|
||||
assert response.headers["content-type"] == "text/markdown; charset=utf-8"
|
||||
assert (
|
||||
normalize_newlines(f"""
|
||||
normalize_newlines(
|
||||
f"""
|
||||
--- memory://test/test-entity {entity1.updated_at.isoformat()} {entity1.checksum[:8]}
|
||||
|
||||
# Test Content
|
||||
@@ -308,7 +315,8 @@ async def test_get_resource_relation(client, project_config, entity_repository,
|
||||
# Related Content
|
||||
- links to [[Test Entity]]
|
||||
|
||||
""".strip())
|
||||
""".strip()
|
||||
)
|
||||
in response.text
|
||||
)
|
||||
|
||||
|
||||
@@ -309,7 +309,9 @@ def test_build_context_with_options(cli_env, setup_test_note):
|
||||
# Check that metadata reflects our options
|
||||
assert context_result["metadata"]["depth"] == 2
|
||||
timeframe = datetime.fromisoformat(context_result["metadata"]["timeframe"])
|
||||
assert datetime.now().astimezone() - timeframe <= timedelta(days=2) # Compare timezone-aware datetimes
|
||||
assert datetime.now().astimezone() - timeframe <= timedelta(
|
||||
days=2
|
||||
) # Compare timezone-aware datetimes
|
||||
|
||||
# Results should include our test note
|
||||
found = False
|
||||
@@ -353,7 +355,11 @@ def test_build_context_string_depth_parameter(cli_env, setup_test_note):
|
||||
)
|
||||
assert result.exit_code == 2 # Typer exits with code 2 for parameter validation errors
|
||||
# Typer should show a usage error for invalid integer
|
||||
assert "invalid" in result.stderr and "is not a valid" in result.stderr and "integer" in result.stderr
|
||||
assert (
|
||||
"invalid" in result.stderr
|
||||
and "is not a valid" in result.stderr
|
||||
and "integer" in result.stderr
|
||||
)
|
||||
|
||||
|
||||
# The get-entity CLI command was removed when tools were refactored
|
||||
|
||||
@@ -87,8 +87,8 @@ def test_project_default_command(mock_reload, mock_run, cli_env):
|
||||
# Patch the os.environ for checking
|
||||
# On Windows, preserve USERPROFILE to allow home directory detection
|
||||
env_vars = {}
|
||||
if os.name == 'nt' and 'USERPROFILE' in os.environ:
|
||||
env_vars['USERPROFILE'] = os.environ['USERPROFILE']
|
||||
if os.name == "nt" and "USERPROFILE" in os.environ:
|
||||
env_vars["USERPROFILE"] = os.environ["USERPROFILE"]
|
||||
|
||||
with patch.dict(os.environ, env_vars, clear=True):
|
||||
# Patch ConfigManager.set_default_project to prevent validation error
|
||||
@@ -192,7 +192,7 @@ def test_project_move_command_uses_permalink(mock_session, mock_call_patch, cli_
|
||||
"""Test that the 'project move' command correctly generates and uses permalink in API call."""
|
||||
# Mock the session to return a current project
|
||||
mock_session.get_current_project.return_value = "current-project"
|
||||
|
||||
|
||||
# Mock successful API response
|
||||
mock_response = MagicMock()
|
||||
mock_response.status_code = 200
|
||||
@@ -202,28 +202,29 @@ def test_project_move_command_uses_permalink(mock_session, mock_call_patch, cli_
|
||||
"default": False,
|
||||
}
|
||||
mock_call_patch.return_value = mock_response
|
||||
|
||||
|
||||
runner = CliRunner()
|
||||
|
||||
|
||||
# Test with a project name that needs normalization (spaces, mixed case)
|
||||
project_name = "Test Project Name"
|
||||
new_path = os.path.join("new", "path", "to", "project")
|
||||
|
||||
|
||||
result = runner.invoke(cli_app, ["project", "move", project_name, new_path])
|
||||
|
||||
|
||||
# Verify command executed successfully
|
||||
assert result.exit_code == 0
|
||||
|
||||
|
||||
# Verify call_patch was called with the correct permalink-formatted project name
|
||||
mock_call_patch.assert_called_once()
|
||||
args, kwargs = mock_call_patch.call_args
|
||||
|
||||
|
||||
# Check the API endpoint uses the normalized permalink
|
||||
expected_endpoint = "/current-project/project/test-project-name"
|
||||
assert args[1] == expected_endpoint # Second argument is the endpoint URL
|
||||
|
||||
|
||||
# Verify the data contains the resolved path (using same normalization as the function)
|
||||
from pathlib import Path
|
||||
|
||||
expected_path = Path(os.path.abspath(os.path.expanduser(new_path))).as_posix()
|
||||
expected_data = {"path": expected_path}
|
||||
assert kwargs["json"] == expected_data
|
||||
|
||||
+1
-1
@@ -53,7 +53,7 @@ def config_home(tmp_path, monkeypatch) -> Path:
|
||||
# Patch HOME environment variable for the duration of the test
|
||||
monkeypatch.setenv("HOME", str(tmp_path))
|
||||
# On Windows, also set USERPROFILE
|
||||
if os.name == 'nt':
|
||||
if os.name == "nt":
|
||||
monkeypatch.setenv("USERPROFILE", str(tmp_path))
|
||||
# Set BASIC_MEMORY_HOME to the test directory
|
||||
monkeypatch.setenv("BASIC_MEMORY_HOME", str(tmp_path / "basic-memory"))
|
||||
|
||||
@@ -13,6 +13,7 @@ constraint with CASCADE DELETE behavior.
|
||||
This test file verifies that the fix works correctly in production databases
|
||||
that have had the migration applied.
|
||||
"""
|
||||
|
||||
from datetime import datetime, timezone
|
||||
|
||||
import pytest
|
||||
@@ -20,37 +21,39 @@ import pytest
|
||||
from basic_memory.services.project_service import ProjectService
|
||||
|
||||
|
||||
#@pytest.mark.skip(reason="Issue #254 not fully resolved yet - foreign key constraint errors still occur")
|
||||
# @pytest.mark.skip(reason="Issue #254 not fully resolved yet - foreign key constraint errors still occur")
|
||||
@pytest.mark.asyncio
|
||||
async def test_issue_254_foreign_key_constraint_fix(project_service: ProjectService, tmp_path):
|
||||
"""Test to verify issue #254 is fixed: project removal with foreign key constraints.
|
||||
|
||||
|
||||
This test reproduces the exact scenario from issue #254:
|
||||
1. Create a project
|
||||
2. Create entities, observations, and relations linked to that project
|
||||
2. Create entities, observations, and relations linked to that project
|
||||
3. Attempt to remove the project
|
||||
4. Verify it succeeds without "FOREIGN KEY constraint failed" errors
|
||||
5. Verify all related data is properly cleaned up via CASCADE DELETE
|
||||
|
||||
|
||||
Once issue #254 is fully fixed, remove the @pytest.mark.skip decorator.
|
||||
"""
|
||||
test_project_name = "issue-254-verification"
|
||||
test_project_path = str(tmp_path / "issue-254-verification")
|
||||
|
||||
|
||||
# Step 1: Create test project
|
||||
await project_service.add_project(test_project_name, test_project_path)
|
||||
project = await project_service.get_project(test_project_name)
|
||||
assert project is not None, "Project should be created successfully"
|
||||
|
||||
|
||||
# Step 2: Create related entities that would cause foreign key constraint issues
|
||||
from basic_memory.repository.entity_repository import EntityRepository
|
||||
from basic_memory.repository.observation_repository import ObservationRepository
|
||||
from basic_memory.repository.relation_repository import RelationRepository
|
||||
|
||||
|
||||
entity_repo = EntityRepository(project_service.repository.session_maker, project_id=project.id)
|
||||
obs_repo = ObservationRepository(project_service.repository.session_maker, project_id=project.id)
|
||||
obs_repo = ObservationRepository(
|
||||
project_service.repository.session_maker, project_id=project.id
|
||||
)
|
||||
rel_repo = RelationRepository(project_service.repository.session_maker, project_id=project.id)
|
||||
|
||||
|
||||
# Create entity
|
||||
entity_data = {
|
||||
"title": "Issue 254 Test Entity",
|
||||
@@ -64,23 +67,23 @@ async def test_issue_254_foreign_key_constraint_fix(project_service: ProjectServ
|
||||
"updated_at": datetime.now(timezone.utc),
|
||||
}
|
||||
entity = await entity_repo.create(entity_data)
|
||||
|
||||
|
||||
# Create observation linked to entity
|
||||
observation_data = {
|
||||
"entity_id": entity.id,
|
||||
"content": "This observation should be cascade deleted",
|
||||
"category": "test"
|
||||
"category": "test",
|
||||
}
|
||||
observation = await obs_repo.create(observation_data)
|
||||
|
||||
|
||||
# Create relation involving the entity
|
||||
relation_data = {
|
||||
"from_id": entity.id,
|
||||
"to_name": "some-other-entity",
|
||||
"relation_type": "relates-to"
|
||||
"relation_type": "relates-to",
|
||||
}
|
||||
relation = await rel_repo.create(relation_data)
|
||||
|
||||
|
||||
# Step 3: Attempt to remove the project
|
||||
# This is where issue #254 manifested - should NOT raise "FOREIGN KEY constraint failed"
|
||||
try:
|
||||
@@ -95,18 +98,18 @@ async def test_issue_254_foreign_key_constraint_fix(project_service: ProjectServ
|
||||
else:
|
||||
# Re-raise unexpected errors
|
||||
raise
|
||||
|
||||
|
||||
# Step 4: Verify project was successfully removed
|
||||
removed_project = await project_service.get_project(test_project_name)
|
||||
assert removed_project is None, "Project should have been removed"
|
||||
|
||||
|
||||
# Step 5: Verify related data was cascade deleted
|
||||
remaining_entity = await entity_repo.find_by_id(entity.id)
|
||||
assert remaining_entity is None, "Entity should have been cascade deleted"
|
||||
|
||||
|
||||
remaining_observation = await obs_repo.find_by_id(observation.id)
|
||||
assert remaining_observation is None, "Observation should have been cascade deleted"
|
||||
|
||||
|
||||
remaining_relation = await rel_repo.find_by_id(relation.id)
|
||||
assert remaining_relation is None, "Relation should have been cascade deleted"
|
||||
|
||||
@@ -114,20 +117,21 @@ async def test_issue_254_foreign_key_constraint_fix(project_service: ProjectServ
|
||||
@pytest.mark.asyncio
|
||||
async def test_issue_254_reproduction(project_service: ProjectService, tmp_path):
|
||||
"""Test that reproduces issue #254 to document the current state.
|
||||
|
||||
|
||||
This test demonstrates the current behavior and will fail until the issue is fixed.
|
||||
It serves as documentation of what the problem was.
|
||||
"""
|
||||
test_project_name = "issue-254-reproduction"
|
||||
test_project_path = str(tmp_path / "issue-254-reproduction")
|
||||
|
||||
|
||||
# Create project and entity
|
||||
await project_service.add_project(test_project_name, test_project_path)
|
||||
project = await project_service.get_project(test_project_name)
|
||||
|
||||
|
||||
from basic_memory.repository.entity_repository import EntityRepository
|
||||
|
||||
entity_repo = EntityRepository(project_service.repository.session_maker, project_id=project.id)
|
||||
|
||||
|
||||
entity_data = {
|
||||
"title": "Reproduction Entity",
|
||||
"entity_type": "note",
|
||||
@@ -140,15 +144,15 @@ async def test_issue_254_reproduction(project_service: ProjectService, tmp_path)
|
||||
"updated_at": datetime.now(timezone.utc),
|
||||
}
|
||||
await entity_repo.create(entity_data)
|
||||
|
||||
|
||||
# This should eventually work without errors once issue #254 is fixed
|
||||
#with pytest.raises(Exception) as exc_info:
|
||||
# with pytest.raises(Exception) as exc_info:
|
||||
await project_service.remove_project(test_project_name)
|
||||
|
||||
|
||||
# Document the current error for tracking
|
||||
# error_message = str(exc_info.value)
|
||||
# assert any(keyword in error_message for keyword in [
|
||||
# "FOREIGN KEY constraint failed",
|
||||
# "constraint",
|
||||
# "integrity"
|
||||
# ]), f"Expected foreign key or integrity constraint error, got: {error_message}"
|
||||
# ]), f"Expected foreign key or integrity constraint error, got: {error_message}"
|
||||
|
||||
@@ -77,7 +77,7 @@ def test_observation_edge_cases():
|
||||
|
||||
def test_observation_excludes_markdown_and_wiki_links():
|
||||
"""Test that markdown links and wiki links are NOT parsed as observations.
|
||||
|
||||
|
||||
This test validates the fix for issue #247 where:
|
||||
- [text](url) markdown links were incorrectly parsed as observations
|
||||
- [[text]] wiki links were incorrectly parsed as observations
|
||||
@@ -85,39 +85,39 @@ def test_observation_excludes_markdown_and_wiki_links():
|
||||
# Test markdown links are NOT observations
|
||||
token = Token("inline", "[Click here](https://example.com)", 0)
|
||||
assert not is_observation(token), "Markdown links should not be parsed as observations"
|
||||
|
||||
token = Token("inline", "[Documentation](./docs/readme.md)", 0)
|
||||
|
||||
token = Token("inline", "[Documentation](./docs/readme.md)", 0)
|
||||
assert not is_observation(token), "Relative markdown links should not be parsed as observations"
|
||||
|
||||
|
||||
token = Token("inline", "[Empty link]()", 0)
|
||||
assert not is_observation(token), "Empty markdown links should not be parsed as observations"
|
||||
|
||||
|
||||
# Test wiki links are NOT observations
|
||||
token = Token("inline", "[[SomeWikiPage]]", 0)
|
||||
assert not is_observation(token), "Wiki links should not be parsed as observations"
|
||||
|
||||
|
||||
token = Token("inline", "[[Multi Word Page]]", 0)
|
||||
assert not is_observation(token), "Multi-word wiki links should not be parsed as observations"
|
||||
|
||||
|
||||
# Test nested brackets are NOT observations
|
||||
token = Token("inline", "[[Nested [[Inner]] Link]]", 0)
|
||||
assert not is_observation(token), "Nested wiki links should not be parsed as observations"
|
||||
|
||||
|
||||
# Test valid observations still work (should return True)
|
||||
token = Token("inline", "[category] This is a valid observation", 0)
|
||||
assert is_observation(token), "Valid observations should still be parsed correctly"
|
||||
|
||||
|
||||
token = Token("inline", "[design] Valid observation #tag", 0)
|
||||
assert is_observation(token), "Valid observations with tags should still work"
|
||||
|
||||
|
||||
token = Token("inline", "Just some text #tag", 0)
|
||||
assert is_observation(token), "Tag-only observations should still work"
|
||||
|
||||
|
||||
# Test edge cases that should NOT be observations
|
||||
token = Token("inline", "[]Empty brackets", 0)
|
||||
assert not is_observation(token), "Empty category brackets should not be observations"
|
||||
|
||||
token = Token("inline", "[category]No space after category", 0)
|
||||
|
||||
token = Token("inline", "[category]No space after category", 0)
|
||||
assert not is_observation(token), "No space after category should not be valid observation"
|
||||
|
||||
|
||||
|
||||
@@ -11,55 +11,55 @@ async def test_write_note_tags_yaml_format(app, project_config):
|
||||
# Create a note with tags using write_note
|
||||
result = await write_note.fn(
|
||||
title="YAML Format Test",
|
||||
folder="test",
|
||||
folder="test",
|
||||
content="Testing YAML tag formatting",
|
||||
tags=["system", "overview", "reference"]
|
||||
tags=["system", "overview", "reference"],
|
||||
)
|
||||
|
||||
|
||||
# Verify the note was created successfully
|
||||
assert "Created note" in result
|
||||
assert "file_path: test/YAML Format Test.md" in result
|
||||
|
||||
|
||||
# Read the file directly to check YAML formatting
|
||||
file_path = project_config.home / "test" / "YAML Format Test.md"
|
||||
content = file_path.read_text(encoding="utf-8")
|
||||
|
||||
|
||||
# Should use YAML list format
|
||||
assert "tags:" in content
|
||||
assert "- system" in content
|
||||
assert "- overview" in content
|
||||
assert "- reference" in content
|
||||
|
||||
|
||||
# Should NOT use JSON array format
|
||||
assert '["system"' not in content
|
||||
assert '"overview"' not in content
|
||||
assert '"reference"]' not in content
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.asyncio
|
||||
async def test_write_note_stringified_json_tags(app, project_config):
|
||||
"""Test that stringified JSON arrays are handled correctly."""
|
||||
# This simulates the issue where AI assistants pass tags as stringified JSON
|
||||
result = await write_note.fn(
|
||||
title="Stringified JSON Test",
|
||||
folder="test",
|
||||
content="Testing stringified JSON tag input",
|
||||
tags='["python", "testing", "json"]' # Stringified JSON array
|
||||
content="Testing stringified JSON tag input",
|
||||
tags='["python", "testing", "json"]', # Stringified JSON array
|
||||
)
|
||||
|
||||
|
||||
# Verify the note was created successfully
|
||||
assert "Created note" in result
|
||||
|
||||
|
||||
# Read the file to check formatting
|
||||
file_path = project_config.home / "test" / "Stringified JSON Test.md"
|
||||
content = file_path.read_text(encoding="utf-8")
|
||||
|
||||
|
||||
# Should properly parse the JSON and format as YAML list
|
||||
assert "tags:" in content
|
||||
assert "- python" in content
|
||||
assert "- testing" in content
|
||||
assert "- json" in content
|
||||
|
||||
|
||||
# Should NOT have the original stringified format issues
|
||||
assert '["python"' not in content
|
||||
assert '"testing"' not in content
|
||||
@@ -73,12 +73,12 @@ async def test_write_note_single_tag_yaml_format(app, project_config):
|
||||
title="Single Tag Test",
|
||||
folder="test",
|
||||
content="Testing single tag formatting",
|
||||
tags=["solo-tag"]
|
||||
tags=["solo-tag"],
|
||||
)
|
||||
|
||||
|
||||
file_path = project_config.home / "test" / "Single Tag Test.md"
|
||||
content = file_path.read_text(encoding="utf-8")
|
||||
|
||||
|
||||
# Single tag should still use list format
|
||||
assert "tags:" in content
|
||||
assert "- solo-tag" in content
|
||||
@@ -88,15 +88,12 @@ async def test_write_note_single_tag_yaml_format(app, project_config):
|
||||
async def test_write_note_no_tags(app, project_config):
|
||||
"""Test that notes without tags work normally."""
|
||||
await write_note.fn(
|
||||
title="No Tags Test",
|
||||
folder="test",
|
||||
content="Testing note without tags",
|
||||
tags=None
|
||||
title="No Tags Test", folder="test", content="Testing note without tags", tags=None
|
||||
)
|
||||
|
||||
|
||||
file_path = project_config.home / "test" / "No Tags Test.md"
|
||||
content = file_path.read_text(encoding="utf-8")
|
||||
|
||||
|
||||
# Should not have tags field in frontmatter
|
||||
assert "tags:" not in content
|
||||
assert "title: No Tags Test" in content
|
||||
@@ -106,20 +103,17 @@ async def test_write_note_no_tags(app, project_config):
|
||||
async def test_write_note_empty_tags_list(app, project_config):
|
||||
"""Test that empty tag lists are handled properly."""
|
||||
await write_note.fn(
|
||||
title="Empty Tags Test",
|
||||
folder="test",
|
||||
content="Testing empty tag list",
|
||||
tags=[]
|
||||
title="Empty Tags Test", folder="test", content="Testing empty tag list", tags=[]
|
||||
)
|
||||
|
||||
|
||||
file_path = project_config.home / "test" / "Empty Tags Test.md"
|
||||
content = file_path.read_text(encoding="utf-8")
|
||||
|
||||
|
||||
# Should not add tags field to frontmatter for empty lists
|
||||
assert "tags:" not in content
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_write_note_update_preserves_yaml_format(app, project_config):
|
||||
"""Test that updating a note preserves the YAML list format."""
|
||||
# First, create the note
|
||||
@@ -127,34 +121,34 @@ async def test_write_note_update_preserves_yaml_format(app, project_config):
|
||||
title="Update Format Test",
|
||||
folder="test",
|
||||
content="Initial content",
|
||||
tags=["initial", "tag"]
|
||||
tags=["initial", "tag"],
|
||||
)
|
||||
|
||||
|
||||
# Then update it with new tags
|
||||
result = await write_note.fn(
|
||||
title="Update Format Test",
|
||||
title="Update Format Test",
|
||||
folder="test",
|
||||
content="Updated content",
|
||||
tags=["updated", "new-tag", "format"]
|
||||
content="Updated content",
|
||||
tags=["updated", "new-tag", "format"],
|
||||
)
|
||||
|
||||
|
||||
# Should be an update, not a new creation
|
||||
assert "Updated note" in result
|
||||
|
||||
|
||||
# Check the file format
|
||||
file_path = project_config.home / "test" / "Update Format Test.md"
|
||||
file_path = project_config.home / "test" / "Update Format Test.md"
|
||||
content = file_path.read_text(encoding="utf-8")
|
||||
|
||||
|
||||
# Should have proper YAML formatting for updated tags
|
||||
assert "tags:" in content
|
||||
assert "- updated" in content
|
||||
assert "- new-tag" in content
|
||||
assert "- format" in content
|
||||
|
||||
|
||||
# Old tags should be gone
|
||||
assert "- initial" not in content
|
||||
assert "- tag" not in content
|
||||
|
||||
|
||||
# Content should be updated
|
||||
assert "Updated content" in content
|
||||
assert "Initial content" not in content
|
||||
@@ -167,15 +161,15 @@ async def test_complex_tags_yaml_format(app, project_config):
|
||||
title="Complex Tags Test",
|
||||
folder="test",
|
||||
content="Testing complex tag formats",
|
||||
tags=["python-3.9", "api_integration", "v2.0", "nested/category", "under_score"]
|
||||
tags=["python-3.9", "api_integration", "v2.0", "nested/category", "under_score"],
|
||||
)
|
||||
|
||||
|
||||
file_path = project_config.home / "test" / "Complex Tags Test.md"
|
||||
content = file_path.read_text(encoding="utf-8")
|
||||
|
||||
|
||||
# All complex tags should format correctly
|
||||
assert "- python-3.9" in content
|
||||
assert "- api_integration" in content
|
||||
assert "- api_integration" in content
|
||||
assert "- v2.0" in content
|
||||
assert "- nested/category" in content
|
||||
assert "- under_score" in content
|
||||
assert "- under_score" in content
|
||||
|
||||
@@ -120,7 +120,7 @@ async def test_build_context_timeframe_formats(client, test_graph):
|
||||
async def test_build_context_string_depth_parameter(client, test_graph):
|
||||
"""Test that build_context handles string depth parameter correctly."""
|
||||
test_url = "memory://test/root"
|
||||
|
||||
|
||||
# Test valid string depth parameter - should either raise ToolError or convert to int
|
||||
try:
|
||||
result = await build_context.fn(url=test_url, depth="2")
|
||||
@@ -130,7 +130,7 @@ async def test_build_context_string_depth_parameter(client, test_graph):
|
||||
except ToolError:
|
||||
# This is also acceptable behavior - type validation should catch it
|
||||
pass
|
||||
|
||||
|
||||
# Test invalid string depth parameter - should raise ToolError
|
||||
with pytest.raises(ToolError):
|
||||
await build_context.fn(url=test_url, depth="invalid")
|
||||
|
||||
@@ -34,7 +34,8 @@ async def test_write_note(app):
|
||||
# Try reading it back via permalink
|
||||
content = await read_note.fn("test/test-note")
|
||||
assert (
|
||||
normalize_newlines(dedent("""
|
||||
normalize_newlines(
|
||||
dedent("""
|
||||
---
|
||||
title: Test Note
|
||||
type: note
|
||||
@@ -46,7 +47,8 @@ async def test_write_note(app):
|
||||
|
||||
# Test
|
||||
This is a test note
|
||||
""").strip())
|
||||
""").strip()
|
||||
)
|
||||
in content
|
||||
)
|
||||
|
||||
@@ -63,7 +65,8 @@ async def test_write_note_no_tags(app):
|
||||
# Should be able to read it back
|
||||
content = await read_note.fn("test/simple-note")
|
||||
assert (
|
||||
normalize_newlines(dedent("""
|
||||
normalize_newlines(
|
||||
dedent("""
|
||||
---
|
||||
title: Simple Note
|
||||
type: note
|
||||
@@ -71,7 +74,8 @@ async def test_write_note_no_tags(app):
|
||||
---
|
||||
|
||||
Just some text
|
||||
""").strip())
|
||||
""").strip()
|
||||
)
|
||||
in content
|
||||
)
|
||||
|
||||
@@ -115,8 +119,9 @@ async def test_write_note_update_existing(app):
|
||||
# Try reading it back
|
||||
content = await read_note.fn("test/test-note")
|
||||
assert (
|
||||
normalize_newlines(dedent(
|
||||
"""
|
||||
normalize_newlines(
|
||||
dedent(
|
||||
"""
|
||||
---
|
||||
title: Test Note
|
||||
type: note
|
||||
@@ -129,7 +134,8 @@ async def test_write_note_update_existing(app):
|
||||
# Test
|
||||
This is an updated note
|
||||
"""
|
||||
).strip())
|
||||
).strip()
|
||||
)
|
||||
== content
|
||||
)
|
||||
|
||||
@@ -394,8 +400,9 @@ async def test_write_note_preserves_content_frontmatter(app):
|
||||
# Try reading it back via permalink
|
||||
content = await read_note.fn("test/test-note")
|
||||
assert (
|
||||
normalize_newlines(dedent(
|
||||
"""
|
||||
normalize_newlines(
|
||||
dedent(
|
||||
"""
|
||||
---
|
||||
title: Test Note
|
||||
type: note
|
||||
@@ -411,7 +418,8 @@ async def test_write_note_preserves_content_frontmatter(app):
|
||||
|
||||
This is a test note
|
||||
"""
|
||||
).strip())
|
||||
).strip()
|
||||
)
|
||||
in content
|
||||
)
|
||||
|
||||
@@ -498,7 +506,8 @@ async def test_write_note_with_custom_entity_type(app):
|
||||
# Verify the entity type is correctly set in the frontmatter
|
||||
content = await read_note.fn("guides/test-guide")
|
||||
assert (
|
||||
normalize_newlines(dedent("""
|
||||
normalize_newlines(
|
||||
dedent("""
|
||||
---
|
||||
title: Test Guide
|
||||
type: guide
|
||||
@@ -510,7 +519,8 @@ async def test_write_note_with_custom_entity_type(app):
|
||||
|
||||
# Guide Content
|
||||
This is a guide
|
||||
""").strip())
|
||||
""").strip()
|
||||
)
|
||||
in content
|
||||
)
|
||||
|
||||
|
||||
@@ -36,7 +36,7 @@ async def second_search_repo(session_maker, second_test_project):
|
||||
@pytest.mark.asyncio
|
||||
async def test_index_item_respects_project_isolation_during_edit():
|
||||
"""Test that index_item() doesn't delete records from other projects during edits.
|
||||
|
||||
|
||||
This test reproduces the critical bug where editing a note in one project
|
||||
would delete search index entries with the same permalink from ALL projects,
|
||||
causing notes to disappear from the search index.
|
||||
@@ -49,11 +49,11 @@ async def test_index_item_respects_project_isolation_during_edit():
|
||||
# Create a separate in-memory database for this test
|
||||
engine = create_async_engine("sqlite+aiosqlite:///:memory:")
|
||||
session_maker = async_sessionmaker(engine, expire_on_commit=False)
|
||||
|
||||
|
||||
# Create the database schema
|
||||
async with engine.begin() as conn:
|
||||
await conn.run_sync(Base.metadata.create_all)
|
||||
|
||||
|
||||
# Create two projects
|
||||
async with db.scoped_session(session_maker) as session:
|
||||
project1 = Project(
|
||||
@@ -61,19 +61,19 @@ async def test_index_item_respects_project_isolation_during_edit():
|
||||
description="First project",
|
||||
path="/project1/path",
|
||||
is_active=True,
|
||||
is_default=True
|
||||
is_default=True,
|
||||
)
|
||||
project2 = Project(
|
||||
name="Project 2",
|
||||
name="Project 2",
|
||||
description="Second project",
|
||||
path="/project2/path",
|
||||
is_active=True,
|
||||
is_default=False
|
||||
is_default=False,
|
||||
)
|
||||
session.add(project1)
|
||||
session.add(project2)
|
||||
await session.flush()
|
||||
|
||||
|
||||
project1_id = project1.id
|
||||
project2_id = project2.id
|
||||
await session.commit()
|
||||
@@ -88,7 +88,7 @@ async def test_index_item_respects_project_isolation_during_edit():
|
||||
# Create two notes with the SAME permalink in different projects
|
||||
# This simulates the same note name/structure across different projects
|
||||
same_permalink = "notes/test-note"
|
||||
|
||||
|
||||
search_row1 = SearchIndexRow(
|
||||
id=1,
|
||||
type=SearchItemType.ENTITY.value,
|
||||
@@ -143,7 +143,7 @@ async def test_index_item_respects_project_isolation_during_edit():
|
||||
content_stems="project 1 content EDITED", # Changed content
|
||||
content_snippet="This is the EDITED content in project 1",
|
||||
permalink=same_permalink,
|
||||
file_path="notes/test_note.md",
|
||||
file_path="notes/test_note.md",
|
||||
entity_id=1,
|
||||
metadata={"entity_type": "note"},
|
||||
created_at=datetime.now(timezone.utc),
|
||||
@@ -175,7 +175,7 @@ async def test_index_item_respects_project_isolation_during_edit():
|
||||
await engine.dispose()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.asyncio
|
||||
async def test_index_item_updates_existing_record_same_project():
|
||||
"""Test that index_item() correctly updates existing records within the same project."""
|
||||
from basic_memory import db
|
||||
@@ -186,11 +186,11 @@ async def test_index_item_updates_existing_record_same_project():
|
||||
# Create a separate in-memory database for this test
|
||||
engine = create_async_engine("sqlite+aiosqlite:///:memory:")
|
||||
session_maker = async_sessionmaker(engine, expire_on_commit=False)
|
||||
|
||||
|
||||
# Create the database schema
|
||||
async with engine.begin() as conn:
|
||||
await conn.run_sync(Base.metadata.create_all)
|
||||
|
||||
|
||||
# Create one project
|
||||
async with db.scoped_session(session_maker) as session:
|
||||
project = Project(
|
||||
@@ -198,7 +198,7 @@ async def test_index_item_updates_existing_record_same_project():
|
||||
description="Test project",
|
||||
path="/test/path",
|
||||
is_active=True,
|
||||
is_default=True
|
||||
is_default=True,
|
||||
)
|
||||
session.add(project)
|
||||
await session.flush()
|
||||
@@ -267,4 +267,4 @@ async def test_index_item_updates_existing_record_same_project():
|
||||
all_results = await repo.search(search_text="My Test Note")
|
||||
assert len(all_results) == 1
|
||||
|
||||
await engine.dispose()
|
||||
await engine.dispose()
|
||||
|
||||
@@ -6,11 +6,11 @@ from datetime import datetime
|
||||
|
||||
from basic_memory.schemas.memory import (
|
||||
EntitySummary,
|
||||
RelationSummary,
|
||||
RelationSummary,
|
||||
ObservationSummary,
|
||||
MemoryMetadata,
|
||||
GraphContext,
|
||||
ContextResult
|
||||
ContextResult,
|
||||
)
|
||||
|
||||
|
||||
@@ -20,18 +20,18 @@ class TestDateTimeSerialization:
|
||||
def test_entity_summary_datetime_serialization(self):
|
||||
"""Test EntitySummary serializes datetime as ISO format string."""
|
||||
test_datetime = datetime(2023, 12, 8, 10, 30, 0)
|
||||
|
||||
|
||||
entity = EntitySummary(
|
||||
permalink="test/entity",
|
||||
title="Test Entity",
|
||||
title="Test Entity",
|
||||
file_path="test/entity.md",
|
||||
created_at=test_datetime
|
||||
created_at=test_datetime,
|
||||
)
|
||||
|
||||
|
||||
# Test model_dump_json() produces ISO format
|
||||
json_str = entity.model_dump_json()
|
||||
data = json.loads(json_str)
|
||||
|
||||
|
||||
assert data["created_at"] == "2023-12-08T10:30:00"
|
||||
assert data["type"] == "entity"
|
||||
assert data["title"] == "Test Entity"
|
||||
@@ -39,21 +39,21 @@ class TestDateTimeSerialization:
|
||||
def test_relation_summary_datetime_serialization(self):
|
||||
"""Test RelationSummary serializes datetime as ISO format string."""
|
||||
test_datetime = datetime(2023, 12, 8, 15, 45, 30)
|
||||
|
||||
|
||||
relation = RelationSummary(
|
||||
title="Test Relation",
|
||||
file_path="test/relation.md",
|
||||
file_path="test/relation.md",
|
||||
permalink="test/relation",
|
||||
relation_type="relates_to",
|
||||
from_entity="entity1",
|
||||
to_entity="entity2",
|
||||
created_at=test_datetime
|
||||
to_entity="entity2",
|
||||
created_at=test_datetime,
|
||||
)
|
||||
|
||||
|
||||
# Test model_dump_json() produces ISO format
|
||||
json_str = relation.model_dump_json()
|
||||
data = json.loads(json_str)
|
||||
|
||||
|
||||
assert data["created_at"] == "2023-12-08T15:45:30"
|
||||
assert data["type"] == "relation"
|
||||
assert data["relation_type"] == "relates_to"
|
||||
@@ -61,20 +61,20 @@ class TestDateTimeSerialization:
|
||||
def test_observation_summary_datetime_serialization(self):
|
||||
"""Test ObservationSummary serializes datetime as ISO format string."""
|
||||
test_datetime = datetime(2023, 12, 8, 20, 15, 45)
|
||||
|
||||
|
||||
observation = ObservationSummary(
|
||||
title="Test Observation",
|
||||
file_path="test/observation.md",
|
||||
permalink="test/observation",
|
||||
permalink="test/observation",
|
||||
category="note",
|
||||
content="Test content",
|
||||
created_at=test_datetime
|
||||
created_at=test_datetime,
|
||||
)
|
||||
|
||||
|
||||
# Test model_dump_json() produces ISO format
|
||||
json_str = observation.model_dump_json()
|
||||
data = json.loads(json_str)
|
||||
|
||||
|
||||
assert data["created_at"] == "2023-12-08T20:15:45"
|
||||
assert data["type"] == "observation"
|
||||
assert data["category"] == "note"
|
||||
@@ -82,18 +82,15 @@ class TestDateTimeSerialization:
|
||||
def test_memory_metadata_datetime_serialization(self):
|
||||
"""Test MemoryMetadata serializes datetime as ISO format string."""
|
||||
test_datetime = datetime(2023, 12, 8, 12, 0, 0)
|
||||
|
||||
|
||||
metadata = MemoryMetadata(
|
||||
depth=2,
|
||||
generated_at=test_datetime,
|
||||
primary_count=5,
|
||||
related_count=3
|
||||
depth=2, generated_at=test_datetime, primary_count=5, related_count=3
|
||||
)
|
||||
|
||||
# Test model_dump_json() produces ISO format
|
||||
|
||||
# Test model_dump_json() produces ISO format
|
||||
json_str = metadata.model_dump_json()
|
||||
data = json.loads(json_str)
|
||||
|
||||
|
||||
assert data["generated_at"] == "2023-12-08T12:00:00"
|
||||
assert data["depth"] == 2
|
||||
assert data["primary_count"] == 5
|
||||
@@ -101,134 +98,117 @@ class TestDateTimeSerialization:
|
||||
def test_context_result_with_datetime_serialization(self):
|
||||
"""Test ContextResult with nested models serializes datetime correctly."""
|
||||
test_datetime = datetime(2023, 12, 8, 9, 30, 15)
|
||||
|
||||
|
||||
entity = EntitySummary(
|
||||
permalink="test/entity",
|
||||
title="Test Entity",
|
||||
file_path="test/entity.md",
|
||||
created_at=test_datetime
|
||||
file_path="test/entity.md",
|
||||
created_at=test_datetime,
|
||||
)
|
||||
|
||||
|
||||
observation = ObservationSummary(
|
||||
title="Test Observation",
|
||||
file_path="test/observation.md",
|
||||
permalink="test/observation",
|
||||
category="note",
|
||||
category="note",
|
||||
content="Test content",
|
||||
created_at=test_datetime
|
||||
created_at=test_datetime,
|
||||
)
|
||||
|
||||
|
||||
context_result = ContextResult(
|
||||
primary_result=entity,
|
||||
observations=[observation],
|
||||
related_results=[]
|
||||
primary_result=entity, observations=[observation], related_results=[]
|
||||
)
|
||||
|
||||
|
||||
# Test model_dump_json() produces ISO format for nested models
|
||||
json_str = context_result.model_dump_json()
|
||||
data = json.loads(json_str)
|
||||
|
||||
|
||||
assert data["primary_result"]["created_at"] == "2023-12-08T09:30:15"
|
||||
assert data["observations"][0]["created_at"] == "2023-12-08T09:30:15"
|
||||
|
||||
def test_graph_context_full_serialization(self):
|
||||
"""Test full GraphContext serialization with all datetime fields."""
|
||||
test_datetime = datetime(2023, 12, 8, 14, 20, 10)
|
||||
|
||||
|
||||
entity = EntitySummary(
|
||||
permalink="test/entity",
|
||||
title="Test Entity",
|
||||
file_path="test/entity.md",
|
||||
created_at=test_datetime
|
||||
created_at=test_datetime,
|
||||
)
|
||||
|
||||
|
||||
metadata = MemoryMetadata(
|
||||
depth=1,
|
||||
generated_at=test_datetime,
|
||||
primary_count=1,
|
||||
related_count=0
|
||||
depth=1, generated_at=test_datetime, primary_count=1, related_count=0
|
||||
)
|
||||
|
||||
context_result = ContextResult(
|
||||
primary_result=entity,
|
||||
observations=[],
|
||||
related_results=[]
|
||||
)
|
||||
|
||||
|
||||
context_result = ContextResult(primary_result=entity, observations=[], related_results=[])
|
||||
|
||||
graph_context = GraphContext(
|
||||
results=[context_result],
|
||||
metadata=metadata,
|
||||
page=1,
|
||||
page_size=10
|
||||
results=[context_result], metadata=metadata, page=1, page_size=10
|
||||
)
|
||||
|
||||
|
||||
# Test full serialization
|
||||
json_str = graph_context.model_dump_json()
|
||||
data = json.loads(json_str)
|
||||
|
||||
|
||||
assert data["metadata"]["generated_at"] == "2023-12-08T14:20:10"
|
||||
assert data["results"][0]["primary_result"]["created_at"] == "2023-12-08T14:20:10"
|
||||
|
||||
def test_datetime_with_microseconds_serialization(self):
|
||||
"""Test datetime with microseconds serializes correctly."""
|
||||
test_datetime = datetime(2023, 12, 8, 10, 30, 0, 123456)
|
||||
|
||||
|
||||
entity = EntitySummary(
|
||||
permalink="test/entity",
|
||||
title="Test Entity",
|
||||
file_path="test/entity.md",
|
||||
created_at=test_datetime
|
||||
created_at=test_datetime,
|
||||
)
|
||||
|
||||
|
||||
json_str = entity.model_dump_json()
|
||||
data = json.loads(json_str)
|
||||
|
||||
|
||||
# Should include microseconds in ISO format
|
||||
assert data["created_at"] == "2023-12-08T10:30:00.123456"
|
||||
|
||||
def test_mcp_schema_validation_compatibility(self):
|
||||
"""Test that serialized datetime format is compatible with MCP schema validation."""
|
||||
test_datetime = datetime(2023, 12, 8, 10, 30, 0)
|
||||
|
||||
|
||||
entity = EntitySummary(
|
||||
permalink="test/entity",
|
||||
title="Test Entity",
|
||||
title="Test Entity",
|
||||
file_path="test/entity.md",
|
||||
created_at=test_datetime
|
||||
created_at=test_datetime,
|
||||
)
|
||||
|
||||
|
||||
# Serialize to JSON
|
||||
json_str = entity.model_dump_json()
|
||||
data = json.loads(json_str)
|
||||
|
||||
|
||||
# Verify the format matches expected MCP "date-time" format
|
||||
datetime_str = data["created_at"]
|
||||
|
||||
|
||||
# Should be parseable back to datetime (ISO format validation)
|
||||
parsed_datetime = datetime.fromisoformat(datetime_str)
|
||||
assert parsed_datetime == test_datetime
|
||||
|
||||
|
||||
# Should match the expected ISO format pattern
|
||||
assert "T" in datetime_str # Contains date-time separator
|
||||
assert len(datetime_str) >= 19 # At least YYYY-MM-DDTHH:MM:SS format
|
||||
|
||||
def test_all_models_have_json_encoders_configured(self):
|
||||
"""Test that all memory schema models have datetime json_encoders configured."""
|
||||
models_to_test = [
|
||||
EntitySummary,
|
||||
RelationSummary,
|
||||
ObservationSummary,
|
||||
MemoryMetadata
|
||||
]
|
||||
|
||||
models_to_test = [EntitySummary, RelationSummary, ObservationSummary, MemoryMetadata]
|
||||
|
||||
for model_class in models_to_test:
|
||||
# Check that ConfigDict with json_encoders is configured
|
||||
assert hasattr(model_class, 'model_config')
|
||||
assert 'json_encoders' in model_class.model_config
|
||||
assert datetime in model_class.model_config['json_encoders']
|
||||
|
||||
assert hasattr(model_class, "model_config")
|
||||
assert "json_encoders" in model_class.model_config
|
||||
assert datetime in model_class.model_config["json_encoders"]
|
||||
|
||||
# Verify the encoder function produces ISO format
|
||||
encoder = model_class.model_config['json_encoders'][datetime]
|
||||
encoder = model_class.model_config["json_encoders"][datetime]
|
||||
test_datetime = datetime(2023, 12, 8, 10, 30, 0)
|
||||
result = encoder(test_datetime)
|
||||
assert result == "2023-12-08T10:30:00"
|
||||
assert result == "2023-12-08T10:30:00"
|
||||
|
||||
@@ -222,14 +222,14 @@ async def test_context_metadata(context_service, test_graph):
|
||||
assert metadata.primary_count > 0
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.asyncio
|
||||
async def test_project_isolation_in_find_related(session_maker):
|
||||
"""Test that find_related respects project boundaries and doesn't leak data."""
|
||||
from basic_memory.repository.entity_repository import EntityRepository
|
||||
from basic_memory.repository.observation_repository import ObservationRepository
|
||||
from basic_memory.repository.search_repository import SearchRepository
|
||||
from basic_memory import db
|
||||
|
||||
|
||||
# Create database session
|
||||
async with db.scoped_session(session_maker) as db_session:
|
||||
# Create two separate projects
|
||||
@@ -238,82 +238,82 @@ async def test_project_isolation_in_find_related(session_maker):
|
||||
db_session.add(project1)
|
||||
db_session.add(project2)
|
||||
await db_session.flush()
|
||||
|
||||
|
||||
# Create entities in project1
|
||||
entity1_p1 = Entity(
|
||||
title="Entity1_P1",
|
||||
entity_type="document",
|
||||
entity_type="document",
|
||||
content_type="text/markdown",
|
||||
project_id=project1.id,
|
||||
permalink="project1/entity1",
|
||||
file_path="project1/entity1.md",
|
||||
created_at=datetime.now(UTC),
|
||||
updated_at=datetime.now(UTC)
|
||||
updated_at=datetime.now(UTC),
|
||||
)
|
||||
entity2_p1 = Entity(
|
||||
title="Entity2_P1",
|
||||
title="Entity2_P1",
|
||||
entity_type="document",
|
||||
content_type="text/markdown",
|
||||
content_type="text/markdown",
|
||||
project_id=project1.id,
|
||||
permalink="project1/entity2",
|
||||
file_path="project1/entity2.md",
|
||||
file_path="project1/entity2.md",
|
||||
created_at=datetime.now(UTC),
|
||||
updated_at=datetime.now(UTC)
|
||||
updated_at=datetime.now(UTC),
|
||||
)
|
||||
|
||||
|
||||
# Create entities in project2
|
||||
entity1_p2 = Entity(
|
||||
title="Entity1_P2",
|
||||
entity_type="document",
|
||||
content_type="text/markdown",
|
||||
project_id=project2.id,
|
||||
project_id=project2.id,
|
||||
permalink="project2/entity1",
|
||||
file_path="project2/entity1.md",
|
||||
created_at=datetime.now(UTC),
|
||||
updated_at=datetime.now(UTC)
|
||||
updated_at=datetime.now(UTC),
|
||||
)
|
||||
|
||||
|
||||
db_session.add_all([entity1_p1, entity2_p1, entity1_p2])
|
||||
await db_session.flush()
|
||||
|
||||
|
||||
# Create relation in project1 (between entities of project1)
|
||||
relation_p1 = Relation(
|
||||
from_id=entity1_p1.id,
|
||||
to_id=entity2_p1.id,
|
||||
to_name="Entity2_P1",
|
||||
relation_type="connects_to"
|
||||
relation_type="connects_to",
|
||||
)
|
||||
db_session.add(relation_p1)
|
||||
await db_session.commit()
|
||||
|
||||
# Create repositories for project1
|
||||
|
||||
# Create repositories for project1
|
||||
search_repo_p1 = SearchRepository(session_maker, project1.id)
|
||||
entity_repo_p1 = EntityRepository(session_maker, project1.id)
|
||||
obs_repo_p1 = ObservationRepository(session_maker, project1.id)
|
||||
context_service_p1 = ContextService(search_repo_p1, entity_repo_p1, obs_repo_p1)
|
||||
|
||||
|
||||
# Create repositories for project2
|
||||
search_repo_p2 = SearchRepository(session_maker, project2.id)
|
||||
entity_repo_p2 = EntityRepository(session_maker, project2.id)
|
||||
entity_repo_p2 = EntityRepository(session_maker, project2.id)
|
||||
obs_repo_p2 = ObservationRepository(session_maker, project2.id)
|
||||
context_service_p2 = ContextService(search_repo_p2, entity_repo_p2, obs_repo_p2)
|
||||
|
||||
|
||||
# Test: find_related for project1 should only return project1 entities
|
||||
type_id_pairs_p1 = [("entity", entity1_p1.id)]
|
||||
related_p1 = await context_service_p1.find_related(type_id_pairs_p1, max_depth=2)
|
||||
|
||||
|
||||
# Verify only project1 entities are returned
|
||||
related_entity_ids = [r.id for r in related_p1 if r.type == "entity"]
|
||||
assert entity2_p1.id in related_entity_ids # Should find connected entity2 in project1
|
||||
assert entity1_p2.id not in related_entity_ids # Should NOT find entity from project2
|
||||
|
||||
|
||||
# Test: find_related for project2 should return empty (no relations)
|
||||
type_id_pairs_p2 = [("entity", entity1_p2.id)]
|
||||
related_p2 = await context_service_p2.find_related(type_id_pairs_p2, max_depth=2)
|
||||
|
||||
|
||||
# Project2 has no relations, so should return empty
|
||||
assert len(related_p2) == 0
|
||||
|
||||
|
||||
# Double-check: verify entities exist in their respective projects
|
||||
assert entity1_p1.project_id == project1.id
|
||||
assert entity2_p1.project_id == project1.id
|
||||
|
||||
@@ -11,10 +11,10 @@ from basic_memory.services.project_service import ProjectService
|
||||
@pytest.mark.asyncio
|
||||
async def test_remove_project_with_related_entities(project_service: ProjectService, tmp_path):
|
||||
"""Test removing a project that has related entities (reproduces issue #254).
|
||||
|
||||
|
||||
This test verifies that projects with related entities (entities, observations, relations)
|
||||
can be properly deleted without foreign key constraint violations.
|
||||
|
||||
|
||||
The bug was caused by missing foreign key constraints with CASCADE DELETE after
|
||||
the project table was recreated in migration 647e7a75e2cd.
|
||||
"""
|
||||
@@ -27,18 +27,21 @@ async def test_remove_project_with_related_entities(project_service: ProjectServ
|
||||
try:
|
||||
# Step 1: Add the test project
|
||||
await project_service.add_project(test_project_name, test_project_path)
|
||||
|
||||
|
||||
# Verify project exists
|
||||
project = await project_service.get_project(test_project_name)
|
||||
assert project is not None
|
||||
|
||||
|
||||
# Step 2: Create related entities for this project
|
||||
from basic_memory.repository.entity_repository import EntityRepository
|
||||
entity_repo = EntityRepository(project_service.repository.session_maker, project_id=project.id)
|
||||
|
||||
|
||||
entity_repo = EntityRepository(
|
||||
project_service.repository.session_maker, project_id=project.id
|
||||
)
|
||||
|
||||
entity_data = {
|
||||
"title": "Test Entity for Deletion",
|
||||
"entity_type": "note",
|
||||
"entity_type": "note",
|
||||
"content_type": "text/markdown",
|
||||
"project_id": project.id,
|
||||
"permalink": "test-deletion-entity",
|
||||
@@ -49,53 +52,59 @@ async def test_remove_project_with_related_entities(project_service: ProjectServ
|
||||
}
|
||||
entity = await entity_repo.create(entity_data)
|
||||
assert entity is not None
|
||||
|
||||
|
||||
# Step 3: Create observations for the entity
|
||||
from basic_memory.repository.observation_repository import ObservationRepository
|
||||
obs_repo = ObservationRepository(project_service.repository.session_maker, project_id=project.id)
|
||||
|
||||
|
||||
obs_repo = ObservationRepository(
|
||||
project_service.repository.session_maker, project_id=project.id
|
||||
)
|
||||
|
||||
observation_data = {
|
||||
"entity_id": entity.id,
|
||||
"content": "This is a test observation",
|
||||
"category": "note"
|
||||
"category": "note",
|
||||
}
|
||||
observation = await obs_repo.create(observation_data)
|
||||
assert observation is not None
|
||||
|
||||
# Step 4: Create relations involving the entity
|
||||
|
||||
# Step 4: Create relations involving the entity
|
||||
from basic_memory.repository.relation_repository import RelationRepository
|
||||
rel_repo = RelationRepository(project_service.repository.session_maker, project_id=project.id)
|
||||
|
||||
|
||||
rel_repo = RelationRepository(
|
||||
project_service.repository.session_maker, project_id=project.id
|
||||
)
|
||||
|
||||
relation_data = {
|
||||
"from_id": entity.id,
|
||||
"to_name": "some-target-entity",
|
||||
"relation_type": "relates-to"
|
||||
"relation_type": "relates-to",
|
||||
}
|
||||
relation = await rel_repo.create(relation_data)
|
||||
assert relation is not None
|
||||
|
||||
|
||||
# Step 5: Attempt to remove the project
|
||||
# This should work with proper cascade delete, or fail with foreign key constraint
|
||||
await project_service.remove_project(test_project_name)
|
||||
|
||||
|
||||
# Step 6: Verify everything was properly deleted
|
||||
|
||||
|
||||
# Project should be gone
|
||||
removed_project = await project_service.get_project(test_project_name)
|
||||
assert removed_project is None, "Project should have been removed"
|
||||
|
||||
|
||||
# Related entities should be cascade deleted
|
||||
remaining_entity = await entity_repo.find_by_id(entity.id)
|
||||
assert remaining_entity is None, "Entity should have been cascade deleted"
|
||||
|
||||
|
||||
# Observations should be cascade deleted
|
||||
remaining_obs = await obs_repo.find_by_id(observation.id)
|
||||
assert remaining_obs is None, "Observation should have been cascade deleted"
|
||||
|
||||
# Relations should be cascade deleted
|
||||
|
||||
# Relations should be cascade deleted
|
||||
remaining_rel = await rel_repo.find_by_id(relation.id)
|
||||
assert remaining_rel is None, "Relation should have been cascade deleted"
|
||||
|
||||
|
||||
except Exception as e:
|
||||
# Check if this is the specific foreign key constraint error from the bug report
|
||||
if "FOREIGN KEY constraint failed" in str(e):
|
||||
@@ -107,7 +116,7 @@ async def test_remove_project_with_related_entities(project_service: ProjectServ
|
||||
else:
|
||||
# Re-raise other unexpected errors
|
||||
raise e
|
||||
|
||||
|
||||
finally:
|
||||
# Clean up - remove project if it still exists
|
||||
if test_project_name in project_service.projects:
|
||||
@@ -119,7 +128,7 @@ async def test_remove_project_with_related_entities(project_service: ProjectServ
|
||||
project_service.config_manager.remove_project(test_project_name)
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
|
||||
project = await project_service.get_project(test_project_name)
|
||||
if project:
|
||||
await project_service.repository.delete(project.id)
|
||||
await project_service.repository.delete(project.id)
|
||||
|
||||
@@ -713,4 +713,4 @@ async def test_synchronize_projects_handles_case_sensitivity_bug(
|
||||
|
||||
db_project = await project_service.repository.get_by_name(name)
|
||||
if db_project:
|
||||
await project_service.repository.delete(db_project.id)
|
||||
await project_service.repository.delete(db_project.id)
|
||||
|
||||
@@ -631,18 +631,14 @@ Testing file timestamps
|
||||
# Check file timestamps
|
||||
file_entity = await entity_service.get_by_permalink("file-dates3")
|
||||
file_stats = file_path.stat()
|
||||
|
||||
|
||||
# Compare using epoch timestamps to handle timezone differences correctly
|
||||
# This ensures we're comparing the actual points in time, not display representations
|
||||
entity_created_epoch = file_entity.created_at.timestamp()
|
||||
entity_updated_epoch = file_entity.updated_at.timestamp()
|
||||
|
||||
assert (
|
||||
abs(entity_created_epoch - file_stats.st_ctime) < 1
|
||||
)
|
||||
assert (
|
||||
abs(entity_updated_epoch - file_stats.st_mtime) < 1
|
||||
) # Allow 1s difference
|
||||
|
||||
assert abs(entity_created_epoch - file_stats.st_ctime) < 1
|
||||
assert abs(entity_updated_epoch - file_stats.st_mtime) < 1 # Allow 1s difference
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
||||
@@ -69,7 +69,9 @@ async def test_handle_changes_empty_set(watch_service, project_config, test_proj
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_handle_vim_atomic_write_delete_still_exists(watch_service, project_config, test_project, sync_service):
|
||||
async def test_handle_vim_atomic_write_delete_still_exists(
|
||||
watch_service, project_config, test_project, sync_service
|
||||
):
|
||||
"""Test vim atomic write scenario: DELETE event but file still exists on disk."""
|
||||
project_dir = project_config.home
|
||||
|
||||
@@ -84,7 +86,7 @@ Initial content for atomic write test
|
||||
"""
|
||||
test_file.write_text(initial_content)
|
||||
await sync_service.sync(project_dir)
|
||||
|
||||
|
||||
# Get initial entity state
|
||||
initial_entity = await sync_service.entity_repository.get_by_file_path("vim_test.md")
|
||||
assert initial_entity is not None
|
||||
@@ -126,21 +128,23 @@ Modified content after atomic write
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_handle_true_deletion_vs_vim_atomic(watch_service, project_config, test_project, sync_service):
|
||||
async def test_handle_true_deletion_vs_vim_atomic(
|
||||
watch_service, project_config, test_project, sync_service
|
||||
):
|
||||
"""Test that true deletions are still handled correctly vs vim atomic writes."""
|
||||
project_dir = project_config.home
|
||||
|
||||
# Create and sync two files
|
||||
atomic_file = project_dir / "atomic_test.md"
|
||||
delete_file = project_dir / "delete_test.md"
|
||||
|
||||
|
||||
content = """---
|
||||
type: note
|
||||
---
|
||||
# Test File
|
||||
Content for testing
|
||||
"""
|
||||
|
||||
|
||||
atomic_file.write_text(content)
|
||||
delete_file.write_text(content)
|
||||
await sync_service.sync(project_dir)
|
||||
@@ -174,16 +178,18 @@ Content for testing
|
||||
events = watch_service.state.recent_events
|
||||
atomic_events = [e for e in events if e.path == "atomic_test.md"]
|
||||
delete_events = [e for e in events if e.path == "delete_test.md"]
|
||||
|
||||
|
||||
assert len(atomic_events) == 1
|
||||
assert atomic_events[0].action == "modified"
|
||||
|
||||
|
||||
assert len(delete_events) == 1
|
||||
assert delete_events[0].action == "deleted"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_handle_vim_atomic_write_markdown_with_relations(watch_service, project_config, test_project, sync_service):
|
||||
async def test_handle_vim_atomic_write_markdown_with_relations(
|
||||
watch_service, project_config, test_project, sync_service
|
||||
):
|
||||
"""Test vim atomic write with markdown files that contain relations."""
|
||||
project_dir = project_config.home
|
||||
|
||||
@@ -241,7 +247,7 @@ This note links to [[Target Note]] multiple times.
|
||||
updated_entity = await sync_service.entity_repository.get_by_file_path("main.md")
|
||||
assert updated_entity is not None
|
||||
assert updated_entity.id == main_entity.id
|
||||
|
||||
|
||||
# Verify relations were processed correctly
|
||||
updated_relations = len(updated_entity.relations)
|
||||
assert updated_relations >= initial_relations # Should have at least as many relations
|
||||
@@ -253,7 +259,9 @@ This note links to [[Target Note]] multiple times.
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_handle_vim_atomic_write_directory_path_ignored(watch_service, project_config, test_project):
|
||||
async def test_handle_vim_atomic_write_directory_path_ignored(
|
||||
watch_service, project_config, test_project
|
||||
):
|
||||
"""Test that directories are properly ignored even in atomic write detection."""
|
||||
project_dir = project_config.home
|
||||
|
||||
|
||||
+26
-21
@@ -6,6 +6,7 @@ import pytest
|
||||
from basic_memory.config import BasicMemoryConfig, ConfigManager
|
||||
from pathlib import Path
|
||||
|
||||
|
||||
class TestBasicMemoryConfig:
|
||||
"""Test BasicMemoryConfig behavior with BASIC_MEMORY_HOME environment variable."""
|
||||
|
||||
@@ -83,84 +84,88 @@ class TestBasicMemoryConfig:
|
||||
|
||||
class TestConfigManager:
|
||||
"""Test ConfigManager functionality."""
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def temp_config_manager(self):
|
||||
"""Create a ConfigManager with temporary config file."""
|
||||
with tempfile.TemporaryDirectory() as temp_dir:
|
||||
temp_path = Path(temp_dir)
|
||||
|
||||
|
||||
# Create a test ConfigManager instance
|
||||
config_manager = ConfigManager()
|
||||
# Override config paths to use temp directory
|
||||
config_manager.config_dir = temp_path / "basic-memory"
|
||||
config_manager.config_file = config_manager.config_dir / "config.yaml"
|
||||
config_manager.config_dir.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
|
||||
# Create initial config with test projects
|
||||
test_config = BasicMemoryConfig(
|
||||
default_project="main",
|
||||
projects={
|
||||
"main": str(temp_path / "main"),
|
||||
"test-project": str(temp_path / "test"),
|
||||
"special-chars": str(temp_path / "special") # This will be the config key for "Special/Chars"
|
||||
}
|
||||
"special-chars": str(
|
||||
temp_path / "special"
|
||||
), # This will be the config key for "Special/Chars"
|
||||
},
|
||||
)
|
||||
config_manager.save_config(test_config)
|
||||
|
||||
|
||||
yield config_manager
|
||||
|
||||
|
||||
def test_set_default_project_with_exact_name_match(self, temp_config_manager):
|
||||
"""Test set_default_project when project name matches config key exactly."""
|
||||
config_manager = temp_config_manager
|
||||
|
||||
|
||||
# Set default to a project that exists with exact name match
|
||||
config_manager.set_default_project("test-project")
|
||||
|
||||
|
||||
# Verify the config was updated
|
||||
config = config_manager.load_config()
|
||||
assert config.default_project == "test-project"
|
||||
|
||||
|
||||
def test_set_default_project_with_permalink_lookup(self, temp_config_manager):
|
||||
"""Test set_default_project when input needs permalink normalization."""
|
||||
config_manager = temp_config_manager
|
||||
|
||||
|
||||
# Simulate a project that was created with special characters
|
||||
# The config key would be the permalink, but user might type the original name
|
||||
|
||||
|
||||
# First add a project with original name that gets normalized
|
||||
config = config_manager.load_config()
|
||||
config.projects["special-chars-project"] = str(Path("/tmp/special"))
|
||||
config_manager.save_config(config)
|
||||
|
||||
|
||||
# Now test setting default using a name that will normalize to the config key
|
||||
config_manager.set_default_project("Special Chars Project") # This should normalize to "special-chars-project"
|
||||
|
||||
config_manager.set_default_project(
|
||||
"Special Chars Project"
|
||||
) # This should normalize to "special-chars-project"
|
||||
|
||||
# Verify the config was updated with the correct config key
|
||||
updated_config = config_manager.load_config()
|
||||
assert updated_config.default_project == "special-chars-project"
|
||||
|
||||
|
||||
def test_set_default_project_uses_canonical_name(self, temp_config_manager):
|
||||
"""Test that set_default_project uses the canonical config key, not user input."""
|
||||
config_manager = temp_config_manager
|
||||
|
||||
|
||||
# Add a project with a config key that differs from user input
|
||||
config = config_manager.load_config()
|
||||
config.projects["my-test-project"] = str(Path("/tmp/mytest"))
|
||||
config_manager.save_config(config)
|
||||
|
||||
|
||||
# Set default using input that will match but is different from config key
|
||||
config_manager.set_default_project("My Test Project") # Should find "my-test-project"
|
||||
|
||||
|
||||
# Verify that the canonical config key is used, not the user input
|
||||
updated_config = config_manager.load_config()
|
||||
assert updated_config.default_project == "my-test-project"
|
||||
# Should NOT be the user input
|
||||
assert updated_config.default_project != "My Test Project"
|
||||
|
||||
|
||||
def test_set_default_project_nonexistent_project(self, temp_config_manager):
|
||||
"""Test set_default_project raises ValueError for nonexistent project."""
|
||||
config_manager = temp_config_manager
|
||||
|
||||
|
||||
with pytest.raises(ValueError, match="Project 'nonexistent' not found"):
|
||||
config_manager.set_default_project("nonexistent")
|
||||
|
||||
@@ -20,7 +20,7 @@ from sqlalchemy.ext.asyncio import create_async_engine, async_sessionmaker
|
||||
|
||||
class ProductionCascadeTest:
|
||||
"""Test cascade delete behavior on production database."""
|
||||
|
||||
|
||||
def __init__(self, db_path: Optional[Path] = None):
|
||||
"""Initialize test with database path."""
|
||||
if db_path is None:
|
||||
@@ -29,94 +29,98 @@ class ProductionCascadeTest:
|
||||
self.db_path = home_dir / ".basic-memory" / "memory.db"
|
||||
else:
|
||||
self.db_path = db_path
|
||||
|
||||
|
||||
# Create backup path
|
||||
self.backup_path = self.db_path.with_suffix('.db.backup')
|
||||
|
||||
self.backup_path = self.db_path.with_suffix(".db.backup")
|
||||
|
||||
self.engine = None
|
||||
self.session_maker = None
|
||||
|
||||
|
||||
async def setup(self):
|
||||
"""Setup database connection."""
|
||||
if not self.db_path.exists():
|
||||
print(f"❌ Production database not found at: {self.db_path}")
|
||||
print("Please ensure Basic Memory has been initialized and the database exists.")
|
||||
sys.exit(1)
|
||||
|
||||
|
||||
print(f"📁 Using database: {self.db_path}")
|
||||
|
||||
|
||||
# Create backup
|
||||
print(f"💾 Creating backup: {self.backup_path}")
|
||||
import shutil
|
||||
|
||||
shutil.copy2(self.db_path, self.backup_path)
|
||||
|
||||
|
||||
# Connect to database
|
||||
db_url = f"sqlite+aiosqlite:///{self.db_path}"
|
||||
self.engine = create_async_engine(db_url, connect_args={"check_same_thread": False})
|
||||
self.session_maker = async_sessionmaker(self.engine, expire_on_commit=False)
|
||||
|
||||
|
||||
async def cleanup(self):
|
||||
"""Cleanup database connection."""
|
||||
if self.engine:
|
||||
await self.engine.dispose()
|
||||
|
||||
|
||||
async def check_foreign_keys_enabled(self) -> bool:
|
||||
"""Check if foreign keys are enabled in this session."""
|
||||
async with self.session_maker() as session:
|
||||
# Enable foreign keys like production does
|
||||
await session.execute(text("PRAGMA foreign_keys=ON"))
|
||||
|
||||
|
||||
result = await session.execute(text("PRAGMA foreign_keys"))
|
||||
fk_enabled = result.fetchone()[0]
|
||||
return bool(fk_enabled)
|
||||
|
||||
|
||||
async def check_schema(self):
|
||||
"""Check current database schema for foreign key constraints."""
|
||||
async with self.session_maker() as session:
|
||||
await session.execute(text("PRAGMA foreign_keys=ON"))
|
||||
|
||||
|
||||
# Check entity table foreign keys
|
||||
result = await session.execute(text("PRAGMA foreign_key_list(entity)"))
|
||||
entity_fks = result.fetchall()
|
||||
|
||||
|
||||
print("🔍 Current entity table foreign key constraints:")
|
||||
for fk in entity_fks:
|
||||
print(f" - Column: {fk[3]} -> {fk[2]}.{fk[4]} (ON DELETE: {fk[6]})")
|
||||
|
||||
|
||||
# Check if CASCADE DELETE is configured
|
||||
has_cascade = any(fk[6] == 'CASCADE' for fk in entity_fks)
|
||||
|
||||
has_cascade = any(fk[6] == "CASCADE" for fk in entity_fks)
|
||||
|
||||
if has_cascade:
|
||||
print("✅ CASCADE DELETE is configured")
|
||||
else:
|
||||
print("❌ CASCADE DELETE is NOT configured (uses NO ACTION)")
|
||||
|
||||
|
||||
return has_cascade
|
||||
|
||||
|
||||
async def create_test_data(self) -> tuple[int, int]:
|
||||
"""Create test project and entity. Returns (project_id, entity_id)."""
|
||||
async with self.session_maker() as session:
|
||||
await session.execute(text("PRAGMA foreign_keys=ON"))
|
||||
|
||||
|
||||
# Create test project
|
||||
project_sql = """
|
||||
INSERT INTO project (name, description, permalink, path, is_active, is_default, created_at, updated_at)
|
||||
VALUES (:name, :description, :permalink, :path, :is_active, :is_default, :created_at, :updated_at)
|
||||
"""
|
||||
now = datetime.now(timezone.utc)
|
||||
|
||||
result = await session.execute(text(project_sql), {
|
||||
"name": "cascade-test-project",
|
||||
"description": "Test project for cascade delete verification",
|
||||
"permalink": "cascade-test-project",
|
||||
"path": "/tmp/cascade-test",
|
||||
"is_active": True,
|
||||
"is_default": False,
|
||||
"created_at": now,
|
||||
"updated_at": now
|
||||
})
|
||||
|
||||
result = await session.execute(
|
||||
text(project_sql),
|
||||
{
|
||||
"name": "cascade-test-project",
|
||||
"description": "Test project for cascade delete verification",
|
||||
"permalink": "cascade-test-project",
|
||||
"path": "/tmp/cascade-test",
|
||||
"is_active": True,
|
||||
"is_default": False,
|
||||
"created_at": now,
|
||||
"updated_at": now,
|
||||
},
|
||||
)
|
||||
project_id = result.lastrowid
|
||||
|
||||
|
||||
# Create test entity linked to project
|
||||
entity_sql = """
|
||||
INSERT INTO entity (title, entity_type, content_type, project_id, permalink, file_path,
|
||||
@@ -124,171 +128,193 @@ class ProductionCascadeTest:
|
||||
VALUES (:title, :entity_type, :content_type, :project_id, :permalink, :file_path,
|
||||
:checksum, :created_at, :updated_at)
|
||||
"""
|
||||
|
||||
result = await session.execute(text(entity_sql), {
|
||||
"title": "Cascade Test Entity",
|
||||
"entity_type": "note",
|
||||
"content_type": "text/markdown",
|
||||
"project_id": project_id,
|
||||
"permalink": "cascade-test-entity",
|
||||
"file_path": "cascade-test-entity.md",
|
||||
"checksum": "test-checksum",
|
||||
"created_at": now,
|
||||
"updated_at": now
|
||||
})
|
||||
|
||||
result = await session.execute(
|
||||
text(entity_sql),
|
||||
{
|
||||
"title": "Cascade Test Entity",
|
||||
"entity_type": "note",
|
||||
"content_type": "text/markdown",
|
||||
"project_id": project_id,
|
||||
"permalink": "cascade-test-entity",
|
||||
"file_path": "cascade-test-entity.md",
|
||||
"checksum": "test-checksum",
|
||||
"created_at": now,
|
||||
"updated_at": now,
|
||||
},
|
||||
)
|
||||
entity_id = result.lastrowid
|
||||
|
||||
|
||||
await session.commit()
|
||||
|
||||
|
||||
print(f"📝 Created test project (ID: {project_id}) and entity (ID: {entity_id})")
|
||||
return project_id, entity_id
|
||||
|
||||
|
||||
async def verify_test_data_exists(self, project_id: int, entity_id: int) -> bool:
|
||||
"""Verify test data exists before deletion."""
|
||||
async with self.session_maker() as session:
|
||||
# Check project exists
|
||||
result = await session.execute(
|
||||
text("SELECT COUNT(*) FROM project WHERE id = :project_id"), {"project_id": project_id}
|
||||
text("SELECT COUNT(*) FROM project WHERE id = :project_id"),
|
||||
{"project_id": project_id},
|
||||
)
|
||||
project_count = result.fetchone()[0]
|
||||
|
||||
|
||||
# Check entity exists
|
||||
result = await session.execute(
|
||||
text("SELECT COUNT(*) FROM entity WHERE id = :entity_id"), {"entity_id": entity_id}
|
||||
)
|
||||
entity_count = result.fetchone()[0]
|
||||
|
||||
|
||||
exists = project_count > 0 and entity_count > 0
|
||||
if exists:
|
||||
print(f"✅ Test data verified: project ({project_count}) and entity ({entity_count}) exist")
|
||||
print(
|
||||
f"✅ Test data verified: project ({project_count}) and entity ({entity_count}) exist"
|
||||
)
|
||||
else:
|
||||
print(f"❌ Test data missing: project ({project_count}) and entity ({entity_count})")
|
||||
|
||||
print(
|
||||
f"❌ Test data missing: project ({project_count}) and entity ({entity_count})"
|
||||
)
|
||||
|
||||
return exists
|
||||
|
||||
|
||||
async def test_cascade_delete(self, project_id: int, entity_id: int) -> bool:
|
||||
"""Test if deleting project cascades to delete entity."""
|
||||
async with self.session_maker() as session:
|
||||
await session.execute(text("PRAGMA foreign_keys=ON"))
|
||||
|
||||
|
||||
try:
|
||||
# Attempt to delete project
|
||||
print(f"🗑️ Attempting to delete project (ID: {project_id})...")
|
||||
|
||||
|
||||
result = await session.execute(
|
||||
text("DELETE FROM project WHERE id = :project_id"), {"project_id": project_id}
|
||||
)
|
||||
|
||||
|
||||
if result.rowcount == 0:
|
||||
print("❌ Project deletion failed - no rows affected")
|
||||
return False
|
||||
|
||||
|
||||
await session.commit()
|
||||
print("✅ Project deletion succeeded")
|
||||
|
||||
|
||||
# Check if entity was cascade deleted
|
||||
result = await session.execute(
|
||||
text("SELECT COUNT(*) FROM entity WHERE id = :entity_id"), {"entity_id": entity_id}
|
||||
text("SELECT COUNT(*) FROM entity WHERE id = :entity_id"),
|
||||
{"entity_id": entity_id},
|
||||
)
|
||||
entity_count = result.fetchone()[0]
|
||||
|
||||
|
||||
if entity_count == 0:
|
||||
print("✅ CASCADE DELETE working: Entity was automatically deleted")
|
||||
return True
|
||||
else:
|
||||
print("❌ CASCADE DELETE NOT working: Entity still exists after project deletion")
|
||||
print(
|
||||
"❌ CASCADE DELETE NOT working: Entity still exists after project deletion"
|
||||
)
|
||||
return False
|
||||
|
||||
|
||||
except Exception as e:
|
||||
await session.rollback()
|
||||
print(f"❌ Project deletion failed with error: {e}")
|
||||
|
||||
|
||||
# Check if it's a foreign key constraint error
|
||||
if "FOREIGN KEY constraint failed" in str(e):
|
||||
print("🔍 This confirms foreign key constraints are enforced but CASCADE DELETE is not configured")
|
||||
|
||||
print(
|
||||
"🔍 This confirms foreign key constraints are enforced but CASCADE DELETE is not configured"
|
||||
)
|
||||
|
||||
return False
|
||||
|
||||
|
||||
async def cleanup_test_data(self, project_id: int, entity_id: int):
|
||||
"""Clean up any remaining test data."""
|
||||
async with self.session_maker() as session:
|
||||
await session.execute(text("PRAGMA foreign_keys=ON"))
|
||||
|
||||
|
||||
try:
|
||||
# Delete entity first (in case cascade didn't work)
|
||||
await session.execute(text("DELETE FROM entity WHERE id = :entity_id"), {"entity_id": entity_id})
|
||||
|
||||
await session.execute(
|
||||
text("DELETE FROM entity WHERE id = :entity_id"), {"entity_id": entity_id}
|
||||
)
|
||||
|
||||
# Delete project
|
||||
await session.execute(text("DELETE FROM project WHERE id = :project_id"), {"project_id": project_id})
|
||||
|
||||
await session.execute(
|
||||
text("DELETE FROM project WHERE id = :project_id"), {"project_id": project_id}
|
||||
)
|
||||
|
||||
await session.commit()
|
||||
print("🧹 Cleaned up any remaining test data")
|
||||
|
||||
|
||||
except Exception as e:
|
||||
print(f"⚠️ Error during cleanup: {e}")
|
||||
await session.rollback()
|
||||
|
||||
|
||||
async def restore_backup(self):
|
||||
"""Restore database from backup."""
|
||||
if self.backup_path.exists():
|
||||
print("🔄 Restoring database from backup...")
|
||||
import shutil
|
||||
|
||||
shutil.copy2(self.backup_path, self.db_path)
|
||||
print("✅ Database restored from backup")
|
||||
|
||||
|
||||
# Remove backup file
|
||||
self.backup_path.unlink()
|
||||
print("🗑️ Backup file removed")
|
||||
else:
|
||||
print("⚠️ No backup file found to restore")
|
||||
|
||||
|
||||
async def run_test(self) -> bool:
|
||||
"""Run the complete cascade delete test."""
|
||||
print("🧪 Production Database CASCADE DELETE Test")
|
||||
print("=" * 50)
|
||||
|
||||
|
||||
try:
|
||||
await self.setup()
|
||||
|
||||
|
||||
# Check if foreign keys are enabled
|
||||
fk_enabled = await self.check_foreign_keys_enabled()
|
||||
print(f"🔐 Foreign keys enabled: {fk_enabled}")
|
||||
|
||||
|
||||
if not fk_enabled:
|
||||
print("❌ Foreign keys are not enabled - this test requires foreign key enforcement")
|
||||
print(
|
||||
"❌ Foreign keys are not enabled - this test requires foreign key enforcement"
|
||||
)
|
||||
return False
|
||||
|
||||
|
||||
# Check current schema
|
||||
has_cascade = await self.check_schema()
|
||||
|
||||
|
||||
# Create test data
|
||||
project_id, entity_id = await self.create_test_data()
|
||||
|
||||
|
||||
# Verify test data exists
|
||||
if not await self.verify_test_data_exists(project_id, entity_id):
|
||||
return False
|
||||
|
||||
|
||||
# Test cascade delete
|
||||
cascade_works = await self.test_cascade_delete(project_id, entity_id)
|
||||
|
||||
|
||||
# Clean up any remaining test data
|
||||
await self.cleanup_test_data(project_id, entity_id)
|
||||
|
||||
|
||||
print("\n" + "=" * 50)
|
||||
print("🧪 TEST RESULTS:")
|
||||
print(f" Schema has CASCADE DELETE: {has_cascade}")
|
||||
print(f" CASCADE DELETE works: {cascade_works}")
|
||||
|
||||
|
||||
if has_cascade and cascade_works:
|
||||
print("✅ PASS: Foreign key constraints are properly configured with CASCADE DELETE")
|
||||
print(
|
||||
"✅ PASS: Foreign key constraints are properly configured with CASCADE DELETE"
|
||||
)
|
||||
elif not has_cascade and not cascade_works:
|
||||
print("❌ FAIL: Foreign key constraints are missing CASCADE DELETE configuration")
|
||||
print("💡 This confirms issue #254 - migration a1b2c3d4e5f6 is needed")
|
||||
else:
|
||||
print("⚠️ MIXED: Unexpected result combination")
|
||||
|
||||
|
||||
return cascade_works
|
||||
|
||||
|
||||
except Exception as e:
|
||||
print(f"💥 Test failed with error: {e}")
|
||||
return False
|
||||
@@ -301,25 +327,27 @@ class ProductionCascadeTest:
|
||||
async def main():
|
||||
"""Main test function."""
|
||||
import argparse
|
||||
|
||||
|
||||
parser = argparse.ArgumentParser(description="Test cascade delete on production database")
|
||||
parser.add_argument("--db-path", type=Path, help="Path to database file (default: ~/.basic-memory/memory.db)")
|
||||
parser.add_argument(
|
||||
"--db-path", type=Path, help="Path to database file (default: ~/.basic-memory/memory.db)"
|
||||
)
|
||||
parser.add_argument("--no-backup", action="store_true", help="Skip creating backup (dangerous)")
|
||||
|
||||
|
||||
args = parser.parse_args()
|
||||
|
||||
|
||||
if args.no_backup:
|
||||
print("⚠️ WARNING: Running without backup!")
|
||||
response = input("Are you sure? Type 'yes' to continue: ")
|
||||
if response.lower() != 'yes':
|
||||
if response.lower() != "yes":
|
||||
print("❌ Aborted")
|
||||
return
|
||||
|
||||
|
||||
test = ProductionCascadeTest(args.db_path)
|
||||
success = await test.run_test()
|
||||
|
||||
|
||||
sys.exit(0 if success else 1)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
asyncio.run(main())
|
||||
asyncio.run(main())
|
||||
|
||||
@@ -11,15 +11,15 @@ def test_tags_formatted_as_yaml_list():
|
||||
post.metadata["title"] = "Test Note"
|
||||
post.metadata["type"] = "note"
|
||||
post.metadata["tags"] = ["system", "overview", "reference"]
|
||||
|
||||
|
||||
result = dump_frontmatter(post)
|
||||
|
||||
|
||||
# Should use YAML list format
|
||||
assert "tags:" in result
|
||||
assert "- system" in result
|
||||
assert "- overview" in result
|
||||
assert "- overview" in result
|
||||
assert "- reference" in result
|
||||
|
||||
|
||||
# Should NOT use JSON array format
|
||||
assert '["system"' not in result
|
||||
assert '"overview"' not in result
|
||||
@@ -31,9 +31,9 @@ def test_empty_tags_list():
|
||||
post = frontmatter.Post("Test content")
|
||||
post.metadata["title"] = "Test Note"
|
||||
post.metadata["tags"] = []
|
||||
|
||||
|
||||
result = dump_frontmatter(post)
|
||||
|
||||
|
||||
# Should have empty list representation
|
||||
assert "tags: []" in result
|
||||
|
||||
@@ -43,21 +43,21 @@ def test_single_tag():
|
||||
post = frontmatter.Post("Test content")
|
||||
post.metadata["title"] = "Test Note"
|
||||
post.metadata["tags"] = ["single-tag"]
|
||||
|
||||
|
||||
result = dump_frontmatter(post)
|
||||
|
||||
|
||||
assert "tags:" in result
|
||||
assert "- single-tag" in result
|
||||
|
||||
|
||||
|
||||
def test_no_tags_metadata():
|
||||
"""Test that posts without tags work normally."""
|
||||
post = frontmatter.Post("Test content")
|
||||
post.metadata["title"] = "Test Note"
|
||||
post.metadata["type"] = "note"
|
||||
|
||||
|
||||
result = dump_frontmatter(post)
|
||||
|
||||
|
||||
assert "title: Test Note" in result
|
||||
assert "type: note" in result
|
||||
assert "tags:" not in result
|
||||
@@ -66,9 +66,9 @@ def test_no_tags_metadata():
|
||||
def test_no_frontmatter():
|
||||
"""Test that posts with no frontmatter just return content."""
|
||||
post = frontmatter.Post("Test content only")
|
||||
|
||||
|
||||
result = dump_frontmatter(post)
|
||||
|
||||
|
||||
assert result == "Test content only"
|
||||
|
||||
|
||||
@@ -77,9 +77,9 @@ def test_complex_tags_with_special_characters():
|
||||
post = frontmatter.Post("Test content")
|
||||
post.metadata["title"] = "Test Note"
|
||||
post.metadata["tags"] = ["python-test", "api_integration", "v2.0", "nested/tag"]
|
||||
|
||||
|
||||
result = dump_frontmatter(post)
|
||||
|
||||
|
||||
assert "- python-test" in result
|
||||
assert "- api_integration" in result
|
||||
assert "- v2.0" in result
|
||||
@@ -91,14 +91,14 @@ def test_tags_order_preserved():
|
||||
post = frontmatter.Post("Test content")
|
||||
post.metadata["title"] = "Test Note"
|
||||
post.metadata["tags"] = ["zebra", "apple", "banana"]
|
||||
|
||||
|
||||
result = dump_frontmatter(post)
|
||||
|
||||
|
||||
# Find the positions of each tag in the output
|
||||
zebra_pos = result.find("- zebra")
|
||||
apple_pos = result.find("- apple")
|
||||
apple_pos = result.find("- apple")
|
||||
banana_pos = result.find("- banana")
|
||||
|
||||
|
||||
# They should appear in the same order as input
|
||||
assert zebra_pos < apple_pos < banana_pos
|
||||
|
||||
@@ -109,14 +109,14 @@ def test_non_tags_lists_also_formatted():
|
||||
post.metadata["title"] = "Test Note"
|
||||
post.metadata["authors"] = ["John Doe", "Jane Smith"]
|
||||
post.metadata["keywords"] = ["AI", "machine learning"]
|
||||
|
||||
|
||||
result = dump_frontmatter(post)
|
||||
|
||||
|
||||
# Authors should be formatted as YAML list
|
||||
assert "authors:" in result
|
||||
assert "- John Doe" in result
|
||||
assert "- Jane Smith" in result
|
||||
|
||||
|
||||
# Keywords should be formatted as YAML list
|
||||
assert "keywords:" in result
|
||||
assert "- AI" in result
|
||||
@@ -127,18 +127,18 @@ def test_mixed_metadata_types():
|
||||
"""Test that mixed metadata types are handled correctly."""
|
||||
post = frontmatter.Post("Test content")
|
||||
post.metadata["title"] = "Test Note"
|
||||
post.metadata["tags"] = ["tag1", "tag2"]
|
||||
post.metadata["tags"] = ["tag1", "tag2"]
|
||||
post.metadata["created"] = "2024-01-01"
|
||||
post.metadata["priority"] = 5
|
||||
post.metadata["draft"] = True
|
||||
|
||||
|
||||
result = dump_frontmatter(post)
|
||||
|
||||
|
||||
# Lists should use YAML format
|
||||
assert "tags:" in result
|
||||
assert "- tag1" in result
|
||||
assert "- tag2" in result
|
||||
|
||||
|
||||
# Other types should be normal
|
||||
assert "title: Test Note" in result
|
||||
assert "created: '2024-01-01'" in result or "created: 2024-01-01" in result
|
||||
@@ -151,13 +151,13 @@ def test_empty_content():
|
||||
post = frontmatter.Post("")
|
||||
post.metadata["title"] = "Empty Note"
|
||||
post.metadata["tags"] = ["empty", "test"]
|
||||
|
||||
|
||||
result = dump_frontmatter(post)
|
||||
|
||||
|
||||
# Should have frontmatter delimiter
|
||||
assert result.startswith("---")
|
||||
assert result.endswith("---\n")
|
||||
|
||||
|
||||
# Should have proper tag formatting
|
||||
assert "- empty" in result
|
||||
assert "- test" in result
|
||||
@@ -169,15 +169,15 @@ def test_roundtrip_compatibility():
|
||||
original_post.metadata["title"] = "Test Note"
|
||||
original_post.metadata["tags"] = ["system", "test", "obsidian"]
|
||||
original_post.metadata["type"] = "note"
|
||||
|
||||
|
||||
# Format with our function
|
||||
formatted = dump_frontmatter(original_post)
|
||||
|
||||
|
||||
# Parse it back
|
||||
parsed_post = frontmatter.loads(formatted)
|
||||
|
||||
|
||||
# Should have same content and metadata
|
||||
assert parsed_post.content == original_post.content
|
||||
assert parsed_post.metadata["title"] == original_post.metadata["title"]
|
||||
assert parsed_post.metadata["tags"] == original_post.metadata["tags"]
|
||||
assert parsed_post.metadata["type"] == original_post.metadata["type"]
|
||||
assert parsed_post.metadata["type"] == original_post.metadata["type"]
|
||||
|
||||
@@ -58,11 +58,11 @@ def test_parse_tags_special_case() -> None:
|
||||
def test_parse_tags_invalid_json() -> None:
|
||||
"""Test that invalid JSON strings fall back to comma-separated parsing."""
|
||||
# Invalid JSON should fall back to comma-separated parsing
|
||||
result = parse_tags('[invalid json')
|
||||
result = parse_tags("[invalid json")
|
||||
assert result == ["[invalid json"] # Treated as single tag
|
||||
|
||||
result = parse_tags('[tag1, tag2]') # Valid bracket format but not JSON
|
||||
|
||||
result = parse_tags("[tag1, tag2]") # Valid bracket format but not JSON
|
||||
assert result == ["[tag1", "tag2]"] # Split by comma
|
||||
|
||||
|
||||
result = parse_tags('["tag1", "tag2"') # Incomplete JSON
|
||||
assert result == ['["tag1"', '"tag2"'] # Fall back to comma separation
|
||||
|
||||
Reference in New Issue
Block a user