Flowcharting / app.py
westland's picture
Fix Streamlit use_container_width compatibility bug
f69ba66
Raw
History Blame Contribute Delete
26 kB
import streamlit as st
import math
import json
import os
import io
from PIL import Image, ImageDraw, ImageFont
# Set page config
st.set_page_config(page_title="B&W Flowchart Sketcher", layout="wide")
# ==========================================
# Geometry & Math Helpers
# ==========================================
def intersect_segments(p1, p2, q1, q2):
x1, y1 = p1
x2, y2 = p2
x3, y3 = q1
x4, y4 = q2
denom = (y4 - y3) * (x2 - x1) - (x4 - x3) * (y2 - y1)
if denom == 0:
return None # Parallel
ua = ((x4 - x3) * (y1 - y3) - (y4 - y3) * (x1 - x3)) / denom
ub = ((x2 - x1) * (y1 - y3) - (y2 - y1) * (x1 - x3)) / denom
if 0 <= ua <= 1 and 0 <= ub <= 1:
return (x1 + ua * (x2 - x1), y1 + ua * (y2 - y1))
return None
def get_boundary_intersection(node, target):
cx, cy = node.x, node.y
tx, ty = target
dx = tx - cx
dy = ty - cy
dist = math.hypot(dx, dy)
if dist == 0:
return (cx, cy)
if node.shape == "circle":
r = min(node.w, node.h) / 2
return (cx + r * dx / dist, cy + r * dy / dist)
elif node.shape == "oval":
a = node.w / 2
b = node.h / 2
t = 1.0 / math.sqrt((dx/a)**2 + (dy/b)**2 + 1e-9)
return (cx + t * dx, cy + t * dy)
else:
if node.shape == "square":
vertices = [
(cx - node.w/2, cy - node.h/2),
(cx + node.w/2, cy - node.h/2),
(cx + node.w/2, cy + node.h/2),
(cx - node.w/2, cy + node.h/2)
]
elif node.shape == "diamond":
vertices = [
(cx, cy - node.h/2),
(cx + node.w/2, cy),
(cx, cy + node.h/2),
(cx - node.w/2, cy)
]
elif node.shape == "inverted triangle":
vertices = [
(cx - node.w/2, cy - node.h/2),
(cx + node.w/2, cy - node.h/2),
(cx, cy + node.h/2)
]
else:
return (cx, cy)
ray_end = (cx + 10000 * dx / dist, cy + 10000 * dy / dist)
for i in range(len(vertices)):
v1 = vertices[i]
v2 = vertices[(i + 1) % len(vertices)]
pt = intersect_segments((cx, cy), ray_end, v1, v2)
if pt:
return pt
return (cx, cy)
def get_quadratic_bezier_points(p0, p1, p2, num_steps=20):
points = []
for i in range(num_steps + 1):
t = i / num_steps
x = (1-t)**2 * p0[0] + 2*(1-t)*t * p1[0] + t**2 * p2[0]
y = (1-t)**2 * p0[1] + 2*(1-t)*t * p1[1] + t**2 * p2[1]
points.append((x, y))
return points
# ==========================================
# Word Wrapping Helper
# ==========================================
def wrap_text_by_width(text, max_width, measure_fn):
if not text:
return ""
words = text.split()
lines = []
current_line = []
for word in words:
test_line = " ".join(current_line + [word])
if measure_fn(test_line) <= max_width:
current_line.append(word)
else:
if current_line:
lines.append(" ".join(current_line))
current_line = [word]
else:
lines.append(word)
current_line = []
if current_line:
lines.append(" ".join(current_line))
return "\n".join(lines)
# ==========================================
# PIL Font Loader
# ==========================================
def get_pil_font(font_name, size, bold=False):
win_font_dir = "C:\\Windows\\Fonts"
paths = []
if bold:
paths.append(os.path.join(win_font_dir, f"{font_name}bd.ttf"))
paths.append(os.path.join(win_font_dir, f"{font_name}b.ttf"))
paths.append(os.path.join(win_font_dir, "arialbd.ttf"))
else:
paths.append(os.path.join(win_font_dir, f"{font_name}.ttf"))
paths.append(os.path.join(win_font_dir, "arial.ttf"))
for p in paths:
if os.path.exists(p):
try:
return ImageFont.truetype(p, size)
except Exception:
pass
return ImageFont.load_default()
# ==========================================
# Model Classes
# ==========================================
class Node:
def __init__(self, designation, shape, description="", x=100.0, y=100.0):
self.id = designation.strip()
self.shape = shape.lower()
self.description = description.strip()
self.x = float(x)
self.y = float(y)
self.w, self.h = self.default_sizes()
def default_sizes(self):
if self.shape == "circle":
return 80.0, 80.0
elif self.shape == "oval":
return 120.0, 60.0
elif self.shape == "square":
return 110.0, 70.0
elif self.shape == "diamond":
return 120.0, 90.0
elif self.shape == "inverted triangle":
return 120.0, 90.0
return 100.0, 60.0
def get_vertices(self):
cx, cy, w, h = self.x, self.y, self.w, self.h
if self.shape == "square":
return [
(cx - w/2, cy - h/2),
(cx + w/2, cy - h/2),
(cx + w/2, cy + h/2),
(cx - w/2, cy + h/2)
]
elif self.shape == "diamond":
return [
(cx, cy - h/2),
(cx + w/2, cy),
(cx, cy + h/2),
(cx - w/2, cy)
]
elif self.shape == "inverted triangle":
return [
(cx - w/2, cy - h/2),
(cx + w/2, cy - h/2),
(cx, cy + h/2)
]
return []
def get_max_text_width(self):
if self.shape == "circle":
return self.w * 0.70
elif self.shape == "oval":
return self.w * 0.75
elif self.shape == "square":
return self.w * 0.82
elif self.shape == "diamond":
return self.w * 0.58
elif self.shape == "inverted triangle":
return self.w * 0.62
return self.w * 0.8
def to_dict(self):
return {
"id": self.id,
"shape": self.shape,
"description": self.description,
"x": self.x,
"y": self.y
}
@classmethod
def from_dict(cls, d):
return cls(d["id"], d["shape"], d.get("description", ""), d["x"], d["y"])
class Edge:
def __init__(self, u, v, style="straight"):
self.u = u.strip()
self.v = v.strip()
self.style = style.lower()
def to_dict(self):
return {
"u": self.u,
"v": self.v,
"style": self.style
}
@classmethod
def from_dict(cls, d):
return cls(d["u"], d["v"], d["style"])
class FlowchartModel:
def __init__(self):
self.nodes = {}
self.edges = []
def add_node(self, designation, shape, description="", x=100.0, y=100.0):
designation = designation.strip()
if not designation:
return False, "Node designation cannot be empty."
node = Node(designation, shape, description, x, y)
self.nodes[designation] = node
return True, node
def delete_node(self, designation):
designation = designation.strip()
if designation in self.nodes:
del self.nodes[designation]
self.edges = [e for e in self.edges if e.u != designation and e.v != designation]
return True
return False
def add_edge_path(self, path_str, style="straight"):
if "->" in path_str or "-->" in path_str:
s = path_str.replace("-->", "->")
parts = [p.strip() for p in s.split("->") if p.strip()]
elif "," in path_str:
parts = [p.strip() for p in path_str.split(",") if p.strip()]
else:
parts = [p.strip() for p in path_str.split() if p.strip()]
if len(parts) < 2:
return False, "Invalid path. Enter at least two nodes (e.g. A -> B)."
missing = [p for p in parts if p not in self.nodes]
if missing:
return False, f"Missing nodes: {', '.join(missing)}. Create them first."
added_count = 0
for i in range(len(parts) - 1):
u, v = parts[i], parts[i+1]
exists = any(e.u == u and e.v == v for e in self.edges)
if not exists:
self.edges.append(Edge(u, v, style))
added_count += 1
return True, f"Added {added_count} edge(s)."
def remove_edge(self, index):
if 0 <= index < len(self.edges):
self.edges.pop(index)
return True
return False
def clear(self):
self.nodes.clear()
self.edges.clear()
def auto_layout(self):
if not self.nodes:
return
adj = {name: [] for name in self.nodes}
in_degree = {name: 0 for name in self.nodes}
for e in self.edges:
if e.u in adj and e.v in adj:
adj[e.u].append(e.v)
in_degree[e.v] += 1
levels = {}
queue = []
for name in self.nodes:
if in_degree[name] == 0:
levels[name] = 0
queue.append(name)
if not queue:
first_node = list(self.nodes.keys())[0]
levels[first_node] = 0
queue.append(first_node)
visited = set()
while queue:
curr = queue.pop(0)
visited.add(curr)
curr_level = levels.get(curr, 0)
for neighbor in adj[curr]:
old_lvl = levels.get(neighbor, -1)
if curr_level + 1 > old_lvl:
levels[neighbor] = curr_level + 1
if neighbor not in visited and neighbor not in queue:
queue.append(neighbor)
for name in self.nodes:
if name not in levels:
levels[name] = 0
nodes_by_level = {}
for name, lvl in levels.items():
nodes_by_level.setdefault(lvl, []).append(name)
canvas_width = 800
level_height = 140
horizontal_spacing = 150
for lvl, lvl_nodes in sorted(nodes_by_level.items()):
lvl_nodes.sort()
num_nodes = len(lvl_nodes)
y = 100 + lvl * level_height
for i, name in enumerate(lvl_nodes):
x = (canvas_width / 2) + (i - (num_nodes - 1) / 2) * horizontal_spacing
node = self.nodes[name]
node.x = x
node.y = y
def save_to_json(self):
data = {
"nodes": [n.to_dict() for n in self.nodes.values()],
"edges": [e.to_dict() for e in self.edges]
}
return json.dumps(data, indent=4)
def load_from_json(self, json_str):
data = json.loads(json_str)
self.clear()
for n_dict in data.get("nodes", []):
node = Node.from_dict(n_dict)
self.nodes[node.id] = node
for e_dict in data.get("edges", []):
self.edges.append(Edge.from_dict(e_dict))
# ==========================================
# PIL Dashed Line Helper
# ==========================================
def draw_dashed_line_pil(draw, points, fill="black", width=2, dash_len=6, gap_len=6):
if len(points) < 2:
return
current_dash_left = dash_len
current_gap_left = 0
drawing = True
for i in range(len(points) - 1):
p1 = points[i]
p2 = points[i+1]
seg_dx = p2[0] - p1[0]
seg_dy = p2[1] - p1[1]
seg_len = math.hypot(seg_dx, seg_dy)
if seg_len == 0:
continue
vx = seg_dx / seg_len
vy = seg_dy / seg_len
dist_moved = 0
curr_pt = p1
while dist_moved < seg_len:
if drawing:
step = min(seg_len - dist_moved, current_dash_left)
next_pt = (curr_pt[0] + vx * step, curr_pt[1] + vy * step)
draw.line([curr_pt, next_pt], fill=fill, width=width)
dist_moved += step
current_dash_left -= step
curr_pt = next_pt
if current_dash_left <= 0:
drawing = False
current_gap_left = gap_len
else:
step = min(seg_len - dist_moved, current_gap_left)
next_pt = (curr_pt[0] + vx * step, curr_pt[1] + vy * step)
dist_moved += step
current_gap_left -= step
curr_pt = next_pt
if current_gap_left <= 0:
drawing = True
current_dash_left = dash_len
# ==========================================
# Flowchart Image Generator
# ==========================================
def render_flowchart_image(model):
if not model.nodes:
img = Image.new("RGB", (400, 300), "white")
return img
# Determine bounds
min_x = min(n.x - n.w/2 for n in model.nodes.values())
max_x = max(n.x + n.w/2 for n in model.nodes.values())
min_y = min(n.y - n.h/2 for n in model.nodes.values())
max_y = max(n.y + n.h/2 for n in model.nodes.values())
margin = 50.0
crop_x1 = min_x - margin
crop_y1 = min_y - margin
crop_x2 = max_x + margin
crop_y2 = max_y + margin
img_w = max(100, int(crop_x2 - crop_x1))
img_h = max(100, int(crop_y2 - crop_y1))
img = Image.new("RGB", (img_w, img_h), "white")
draw = ImageDraw.Draw(img)
def get_pt(cx, cy):
return int(cx - crop_x1), int(cy - crop_y1)
font_bold = get_pil_font("segoeui", 10, bold=True)
font_reg = get_pil_font("segoeui", 9, bold=False)
if hasattr(draw, 'textlength'):
measure_fn = lambda txt: draw.textlength(txt, font=font_reg)
else:
measure_fn = lambda txt: font_reg.getsize(txt)[0]
def draw_pil_arrowhead(p_end, p_from):
xe, ye = get_pt(p_end[0], p_end[1])
xf, yf = get_pt(p_from[0], p_from[1])
dx = xe - xf
dy = ye - yf
dist = math.hypot(dx, dy)
if dist == 0:
return
ux = dx / dist
uy = dy / dist
arrow_l = 12
arrow_w = 8
bx = xe - ux * arrow_l
by = ye - uy * arrow_l
px = -uy * (arrow_w / 2)
py = ux * (arrow_w / 2)
pts = [(xe, ye), (bx + px, by + py), (bx - px, by - py)]
draw.polygon(pts, fill="black", outline="black")
def draw_centered_text_pil(xy, text, font, fill="black"):
lines = text.split("\n")
try:
bbox = font.getbbox("Abgqp")
line_h = bbox[3] - bbox[1]
except:
line_h = font.getsize("Abgqp")[1]
line_widths = []
for line in lines:
if hasattr(draw, 'textlength'):
w = draw.textlength(line, font=font)
else:
w = font.getsize(line)[0]
line_widths.append(w)
total_height = line_h * len(lines) + 4 * (len(lines) - 1)
curr_y = xy[1] - total_height / 2
for i, line in enumerate(lines):
w = line_widths[i]
draw.text((xy[0] - w/2, curr_y), line, font=font, fill=fill)
curr_y += line_h + 4
# 1. Draw Edges
for edge in model.edges:
u_node = model.nodes.get(edge.u)
v_node = model.nodes.get(edge.v)
if not u_node or not v_node:
continue
if edge.style == "curved":
dx = v_node.x - u_node.x
dy = v_node.y - u_node.y
dist = math.hypot(dx, dy)
mx = (u_node.x + v_node.x) / 2
my = (u_node.y + v_node.y) / 2
if dist > 0:
nx = -dy / dist
ny = dx / dist
p_ctrl = (mx + nx * 45, my + ny * 45)
else:
p_ctrl = (mx, my + 45)
p_start = get_boundary_intersection(u_node, p_ctrl)
p_end = get_boundary_intersection(v_node, p_ctrl)
bezier_pts = get_quadratic_bezier_points(p_start, p_ctrl, p_end, 30)
mapped_pts = [get_pt(px, py) for px, py in bezier_pts]
if edge.style == "dotted":
draw_dashed_line_pil(draw, mapped_pts, fill="black", width=2, dash_len=4, gap_len=4)
else:
draw.line(mapped_pts, fill="black", width=2)
draw_pil_arrowhead(p_end, p_ctrl)
else:
p_start = get_boundary_intersection(u_node, (v_node.x, v_node.y))
p_end = get_boundary_intersection(v_node, (u_node.x, u_node.y))
p_start_m = get_pt(p_start[0], p_start[1])
p_end_m = get_pt(p_end[0], p_end[1])
if edge.style == "dotted":
steps = int(math.hypot(p_end_m[0] - p_start_m[0], p_end_m[1] - p_start_m[1]) / 6)
if steps < 2: steps = 2
pts = []
for s in range(steps + 1):
t = s / steps
pts.append((p_start_m[0] + t * (p_end_m[0] - p_start_m[0]), p_start_m[1] + t * (p_end_m[1] - p_start_m[1])))
draw_dashed_line_pil(draw, pts, fill="black", width=2, dash_len=4, gap_len=4)
else:
draw.line([p_start_m, p_end_m], fill="black", width=2)
draw_pil_arrowhead(p_end, p_start)
# 2. Draw Nodes
for node in model.nodes.values():
cx, cy, w, h = node.x, node.y, node.w, node.h
px, py = get_pt(cx, cy)
# Shadow
sh = 4
shadow_color = (233, 236, 239)
if node.shape in ("circle", "oval"):
draw.ellipse([px - w/2 + sh, py - h/2 + sh, px + w/2 + sh, py + h/2 + sh], fill=shadow_color, outline=None)
elif node.shape == "square":
draw.rectangle([px - w/2 + sh, py - h/2 + sh, px + w/2 + sh, py + h/2 + sh], fill=shadow_color, outline=None)
else:
sh_vertices = [(vx - crop_x1 + sh, vy - crop_y1 + sh) for vx, vy in node.get_vertices()]
draw.polygon(sh_vertices, fill=shadow_color, outline=None)
# Shape
if node.shape in ("circle", "oval"):
draw.ellipse([px - w/2, py - h/2, px + w/2, py + h/2], fill="white", outline="black", width=2)
elif node.shape == "square":
draw.rectangle([px - w/2, py - h/2, px + w/2, py + h/2], fill="white", outline="black", width=2)
else:
sh_vertices = [(vx - crop_x1, vy - crop_y1) for vx, vy in node.get_vertices()]
draw.polygon(sh_vertices, fill="white")
draw.line(sh_vertices + [sh_vertices[0]], fill="black", width=2, joint="curve")
# Text
max_w = node.get_max_text_width()
if node.description:
desc_wrapped = wrap_text_by_width(node.description, max_w, measure_fn)
draw_centered_text_pil((px, py - 12), node.id, font_bold, fill="black")
draw_centered_text_pil((px, py + 10), desc_wrapped, font_reg, fill=(73, 80, 87))
else:
draw_centered_text_pil((px, py), node.id, font_bold, fill="black")
return img
# ==========================================
# Streamlit App Logic
# ==========================================
st.title("📊 B&W Flowchart Sketcher")
# Initialize Model in Session State
if "model" not in st.session_state:
model = FlowchartModel()
# Load defaults
model.add_node("Start", "Circle", "Start", 400.0, 80.0)
model.add_node("Step 1", "Square", "Process Inputs", 400.0, 200.0)
model.add_node("Check", "Diamond", "Valid?", 400.0, 330.0)
model.add_node("Error", "Inverted Triangle", "Log Error", 220.0, 330.0)
model.add_node("End", "Oval", "Success Finish", 400.0, 460.0)
model.add_edge_path("Start -> Step 1", "straight")
model.add_edge_path("Step 1 -> Check", "straight")
model.add_edge_path("Check -> Error", "straight")
model.add_edge_path("Check -> End", "straight")
model.add_edge_path("Error -> Start", "curved")
st.session_state.model = model
else:
model = st.session_state.model
# Split Layout
col_ctrl, col_canvas = st.columns([1, 2])
with col_ctrl:
st.header("Controls")
# --- Node Management ---
with st.expander("Node Manager", expanded=True):
node_id = st.text_input("Designation / Label (Short ID)", key="node_id")
shape = st.selectbox("Icon Shape", ["Square", "Circle", "Oval", "Diamond", "Inverted Triangle"], key="shape")
desc = st.text_input("Description (Inside Shape)", key="desc")
col_node_btns = st.columns(2)
with col_node_btns[0]:
if st.button("Add/Save Node"):
if node_id:
is_update = node_id in model.nodes
if is_update:
n = model.nodes[node_id]
n.shape = shape.lower()
n.description = desc
n.w, n.h = n.default_sizes()
else:
model.add_node(node_id, shape, desc, x=300.0 + 50.0 * len(model.nodes), y=200.0)
st.rerun()
else:
st.error("Short ID designation required.")
with col_node_btns[1]:
if st.button("Delete Node"):
if node_id:
if model.delete_node(node_id):
st.rerun()
else:
st.error(f"Node '{node_id}' not found.")
else:
st.error("Enter a Short ID designation to delete.")
# --- Edge Management ---
with st.expander("Edge Manager", expanded=True):
edge_path = st.text_input("Path (e.g., A -> B -> C)", key="edge_path")
edge_style = st.selectbox("Connector Line Style", ["Straight", "Curved", "Dotted"], key="edge_style")
if st.button("Add Edges"):
if edge_path:
success, msg = model.add_edge_path(edge_path, edge_style)
if success:
st.rerun()
else:
st.error(msg)
else:
st.error("Enter an edge path.")
# --- Manual Node Positioning ---
with st.expander("Manual Node Positioning", expanded=False):
if model.nodes:
selected_pos_node = st.selectbox("Select Node to Position", list(model.nodes.keys()))
node_to_move = model.nodes[selected_pos_node]
new_x = st.slider("X Coordinate", 0.0, 1000.0, float(node_to_move.x), step=10.0)
new_y = st.slider("Y Coordinate", 0.0, 1000.0, float(node_to_move.y), step=10.0)
if new_x != node_to_move.x or new_y != node_to_move.y:
node_to_move.x = new_x
node_to_move.y = new_y
st.rerun()
else:
st.info("Create nodes to enable manual positioning.")
# --- Lists and Deletions ---
with st.expander("Active Lists"):
st.subheader("Nodes")
if model.nodes:
for k, v in model.nodes.items():
st.text(f"• {k} ({v.shape.title()}) - {v.description}")
else:
st.text("No nodes.")
st.subheader("Edges")
if model.edges:
edge_labels = [f"{e.u} -> {e.v} ({e.style.title()})" for e in model.edges]
del_edge_idx = st.selectbox("Select Edge to Delete", range(len(model.edges)), format_func=lambda i: edge_labels[i])
if st.button("Delete Selected Edge"):
model.remove_edge(del_edge_idx)
st.rerun()
else:
st.text("No edges.")
# --- Actions ---
with st.expander("Global Actions", expanded=True):
col_actions = st.columns(2)
with col_actions[0]:
if st.button("Auto-Layout"):
model.auto_layout()
st.rerun()
with col_actions[1]:
if st.button("Clear All"):
model.clear()
st.rerun()
# Save and Load
st.subheader("Save / Load Project")
json_data = model.save_to_json()
st.download_button(
label="Download JSON Project File",
data=json_data,
file_name="flowchart_project.json",
mime="application/json"
)
uploaded_file = st.file_uploader("Upload JSON Project", type="json")
if uploaded_file is not None:
try:
json_str = uploaded_file.read().decode("utf-8")
model.load_from_json(json_str)
st.success("Project loaded successfully!")
st.rerun()
except Exception as e:
st.error(f"Error loading project: {e}")
with col_canvas:
st.header("Sketch Canvas")
# Render flowchart B&W image
flowchart_img = render_flowchart_image(model)
# Display in browser
buffer = io.BytesIO()
flowchart_img.save(buffer, format="PNG")
img_bytes = buffer.getvalue()
st.image(img_bytes, caption="Clean B&W Flowchart (Auto-cropped)", use_column_width=False)
# Download Button
st.download_button(
label="📥 Download Flowchart Image (PNG)",
data=img_bytes,
file_name="flowchart_export.png",
mime="image/png"
)