Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 1 addition & 1 deletion backend/app/__init__.py
Original file line number Diff line number Diff line change
@@ -1 +1 @@
"""AI-Workflow backend application package."""
"""Berry AI Studio backend application package."""
122 changes: 109 additions & 13 deletions backend/app/core/cache.py
Original file line number Diff line number Diff line change
@@ -1,53 +1,149 @@
"""Deterministic hashing and output result caching engine."""
"""Deterministic hashing and output result caching engine with SQLite persistence."""

import hashlib
import json
from typing import Any, Dict, Optional
from datetime import datetime, timezone
from typing import Any, Dict, List, Optional, Tuple

from app.storage.db import DatabaseManager, db_manager


def compute_content_hash(value: Any) -> str:
"""Compute a deterministic SHA-256 hash of output content."""
if value is None:
return hashlib.sha256(b"null").hexdigest()
if isinstance(value, (int, float, bool)):
return hashlib.sha256(str(value).encode("utf-8")).hexdigest()
if isinstance(value, str):
return hashlib.sha256(value.encode("utf-8")).hexdigest()
if isinstance(value, (dict, list)):
serialized = json.dumps(value, sort_keys=True, default=str)
return hashlib.sha256(serialized.encode("utf-8")).hexdigest()
return hashlib.sha256(str(value).encode("utf-8")).hexdigest()


def compute_semantic_node_hash(
node_type: str,
params: Dict[str, Any],
input_bindings: List[Tuple[str, str, str]], # (target_handle, upstream_content_hash, source_handle)
) -> str:
"""
Compute a deterministic port-aware SHA-256 hash representing a node's exact semantic state.
NodeHash = SHA256(NodeType + SerializedCanonicalParams + OrderedPortBindings)

OrderedPortBindings strictly associates input ports to upstream outputs:
target_handle:source_handle:content_hash
"""
hasher = hashlib.sha256()
hasher.update(node_type.encode("utf-8"))

# Canonicalize params: exclude transient secrets from hash
clean_params = {k: v for k, v in params.items() if not k.lower().endswith("key")}
params_str = json.dumps(clean_params, sort_keys=True, default=str)
hasher.update(params_str.encode("utf-8"))

# Deterministic binding sort by target_handle
sorted_bindings = sorted(input_bindings, key=lambda b: (b[0], b[2]))
for target_handle, content_hash, source_handle in sorted_bindings:
binding_repr = f"|in:{target_handle}->out:{source_handle}#{content_hash}"
hasher.update(binding_repr.encode("utf-8"))

return hasher.hexdigest()


def compute_node_hash(node_type: str, params: Dict[str, Any], parent_hashes: list[str]) -> str:
"""
Compute a deterministic SHA-256 hash representing a node's exact state.
NodeHash = SHA256(NodeType + SerializedParams + SortedParentHashes)
Legacy helper: Compute SHA-256 representing a node's state with parent hashes.
Kept for backwards compatibility with baseline tests.
"""
hasher = hashlib.sha256()
hasher.update(node_type.encode("utf-8"))

# Serialize parameters with sorted keys for deterministic encoding
params_str = json.dumps(params, sort_keys=True, default=str)
hasher.update(params_str.encode("utf-8"))

# Concatenate upstream parent hashes in deterministic sorted order
for parent_hash in sorted(parent_hashes):
hasher.update(parent_hash.encode("utf-8"))

return hasher.hexdigest()


class CacheStore:
"""Thread-safe in-memory cache store for node execution outputs."""
"""Thread-safe cache store with in-memory cache and SQLite persistence."""

def __init__(self) -> None:
def __init__(self, manager: Optional[DatabaseManager] = None) -> None:
self._store: Dict[str, Dict[str, Any]] = {}
self.manager = manager or db_manager

def get(self, node_hash: str) -> Optional[Dict[str, Any]]:
"""Retrieve cached output data by node hash."""
"""Retrieve cached output data by node hash synchronously from memory."""
return self._store.get(node_hash)

def set(self, node_hash: str, output: Dict[str, Any]) -> None:
"""Store output data for a node hash."""
"""Store output data for a node hash in memory."""
self._store[node_hash] = output

def has(self, node_hash: str) -> bool:
"""Check if output for node hash is cached."""
"""Check if output for node hash is cached in memory."""
return node_hash in self._store

async def get_async(self, node_hash: str) -> Optional[Dict[str, Any]]:
"""Retrieve cached output, checking memory first then SQLite."""
if node_hash in self._store:
return self._store[node_hash]

try:
conn = await self.manager.get_connection()
async with conn.execute(
"SELECT output_json FROM cache_entries WHERE node_hash = ?",
(node_hash,),
) as cursor:
row = await cursor.fetchone()
if row:
output = json.loads(row["output_json"])
self._store[node_hash] = output
return output
except Exception:
pass
return None

async def set_async(self, node_hash: str, output: Dict[str, Any]) -> None:
"""Store output data in both memory and SQLite."""
self._store[node_hash] = output
try:
conn = await self.manager.get_connection()
now = datetime.now(timezone.utc).isoformat()
await conn.execute(
"""
INSERT OR REPLACE INTO cache_entries (node_hash, output_json, asset_ids_json, created_at)
VALUES (?, ?, ?, ?)
""",
(node_hash, json.dumps(output), json.dumps([]), now),
)
await conn.commit()
except Exception:
pass

async def has_async(self, node_hash: str) -> bool:
"""Check if output exists in memory or SQLite."""
return (await self.get_async(node_hash)) is not None

def clear(self) -> None:
"""Purge all cached results."""
"""Purge in-memory cached results."""
self._store.clear()

async def clear_all_async(self) -> None:
"""Purge both in-memory and persistent SQLite cache."""
self._store.clear()
try:
conn = await self.manager.get_connection()
await conn.execute("DELETE FROM cache_entries")
await conn.commit()
except Exception:
pass

def size(self) -> int:
"""Return number of cached node states."""
"""Return number of in-memory cached node states."""
return len(self._store)


Expand Down
11 changes: 11 additions & 0 deletions backend/app/core/dag.py
Original file line number Diff line number Diff line change
Expand Up @@ -77,6 +77,17 @@ def _get_ancestors(self, node_id: str) -> Set[str]:
queue.extend(self.parents[curr])
return ancestors

def get_descendants(self, node_id: str) -> Set[str]:
"""Recursively gather all downstream descendant nodes."""
descendants: Set[str] = set()
queue = deque(self.adjacency[node_id])
while queue:
curr = queue.popleft()
if curr not in descendants:
descendants.add(curr)
queue.extend(self.adjacency[curr])
return descendants

def get_parent_ids(self, node_id: str) -> List[str]:
"""Return direct upstream parent node IDs."""
return self.parents.get(node_id, [])
Loading
Loading