Skip to content

Commit d78f867

Browse files
Preserve TreeView state during keyed reorders
1 parent f6e428c commit d78f867

4 files changed

Lines changed: 626 additions & 1 deletion

File tree

collagraph/renderers/pyside/objects/itemmodel.py

Lines changed: 87 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -1,8 +1,76 @@
1-
from PySide6.QtGui import QStandardItemModel
1+
from PySide6.QtCore import QItemSelectionModel
2+
from PySide6.QtGui import QStandardItem, QStandardItemModel
3+
from PySide6.QtWidgets import QListView, QTableView, QTreeView
24

35
from ... import PySideRenderer
46

57

8+
def _get_view_from_model(model: QStandardItemModel):
9+
parent = model.parent()
10+
if isinstance(parent, (QTreeView, QListView, QTableView)):
11+
return parent
12+
return None
13+
14+
15+
def _snapshot_item_state(
16+
item: QStandardItem,
17+
model: QStandardItemModel,
18+
view,
19+
) -> list[tuple[QStandardItem, bool, bool]]:
20+
snapshot: list[tuple[QStandardItem, bool, bool]] = []
21+
22+
def walk(it: QStandardItem) -> None:
23+
index = model.indexFromItem(it)
24+
selected = view.selectionModel().isSelected(index)
25+
expanded = isinstance(view, QTreeView) and view.isExpanded(index)
26+
snapshot.append((it, expanded, selected))
27+
for row in range(it.rowCount()):
28+
child = it.child(row)
29+
if child:
30+
walk(child)
31+
32+
walk(item)
33+
return snapshot
34+
35+
36+
def _apply_item_selection(
37+
model: QStandardItemModel,
38+
view,
39+
item: QStandardItem,
40+
selected: bool,
41+
) -> None:
42+
selection_model = view.selectionModel()
43+
index = model.indexFromItem(item)
44+
flags = QItemSelectionModel.SelectionFlag.Rows
45+
if selected:
46+
flags |= QItemSelectionModel.SelectionFlag.Select
47+
else:
48+
flags |= QItemSelectionModel.SelectionFlag.Deselect
49+
selection_model.select(index, flags)
50+
51+
52+
def _apply_item_expanded(
53+
model: QStandardItemModel,
54+
view,
55+
item: QStandardItem,
56+
expanded: bool,
57+
) -> None:
58+
if isinstance(view, QTreeView):
59+
index = model.indexFromItem(item)
60+
view.setExpanded(index, expanded)
61+
62+
63+
def _restore_item_state(
64+
snapshot: list[tuple[QStandardItem, bool, bool]],
65+
model: QStandardItemModel,
66+
view,
67+
) -> None:
68+
for item, expanded, _selected in snapshot:
69+
_apply_item_expanded(model, view, item, expanded)
70+
for item, _expanded, selected in snapshot:
71+
_apply_item_selection(model, view, item, selected)
72+
73+
674
@PySideRenderer.register_insert(QStandardItemModel)
775
def insert(self, el, anchor=None):
876
if isinstance(self, QStandardItemModel):
@@ -16,11 +84,29 @@ def insert(self, el, anchor=None):
1684
self.insertRow(index.row(), el)
1785
else:
1886
self.appendRow(el)
87+
88+
view = _get_view_from_model(self)
89+
if view:
90+
if hasattr(el, "_state_snapshot"):
91+
_restore_item_state(el._state_snapshot, self, view)
92+
del el._state_snapshot
93+
94+
if hasattr(el, "_expanded"):
95+
_apply_item_expanded(self, view, el, el._expanded)
96+
delattr(el, "_expanded")
97+
98+
if hasattr(el, "_selected"):
99+
_apply_item_selection(self, view, el, el._selected)
100+
delattr(el, "_selected")
19101
else:
20102
raise NotImplementedError(type(self).__name__)
21103

22104

23105
@PySideRenderer.register_remove(QStandardItemModel)
24106
def remove(self, el):
107+
view = _get_view_from_model(self)
108+
if view:
109+
el._state_snapshot = _snapshot_item_state(el, self, view)
110+
25111
index = self.indexFromItem(el)
26112
self.takeRow(index.row())

collagraph/renderers/pyside/objects/standarditem.py

Lines changed: 113 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1,9 +1,98 @@
1+
from PySide6.QtCore import QItemSelectionModel
12
from PySide6.QtGui import QStandardItem
3+
from PySide6.QtWidgets import QListView, QTableView, QTreeView
24

35
from ... import PySideRenderer
46
from .qobject import set_attribute as qobject_set_attribute
57

68

9+
def _get_view_from_item(item: QStandardItem):
10+
model = item.model()
11+
if not model:
12+
return None
13+
14+
parent = model.parent()
15+
if isinstance(parent, (QTreeView, QListView, QTableView)):
16+
return parent
17+
return None
18+
19+
20+
def _snapshot_item_state(item: QStandardItem):
21+
model = item.model()
22+
view = _get_view_from_item(item)
23+
if not (model and view):
24+
return None
25+
26+
snapshot: list[tuple[QStandardItem, bool, bool]] = []
27+
28+
def walk(it: QStandardItem) -> None:
29+
index = model.indexFromItem(it)
30+
selected = view.selectionModel().isSelected(index)
31+
expanded = isinstance(view, QTreeView) and view.isExpanded(index)
32+
snapshot.append((it, expanded, selected))
33+
for row in range(it.rowCount()):
34+
child = it.child(row)
35+
if child:
36+
walk(child)
37+
38+
walk(item)
39+
return snapshot
40+
41+
42+
def _apply_item_selection(item: QStandardItem, selected: bool) -> None:
43+
model = item.model()
44+
view = _get_view_from_item(item)
45+
if not (model and view):
46+
item._selected = selected
47+
return
48+
49+
selection_model = view.selectionModel()
50+
index = model.indexFromItem(item)
51+
flags = QItemSelectionModel.SelectionFlag.Rows
52+
if selected:
53+
flags |= QItemSelectionModel.SelectionFlag.Select
54+
else:
55+
flags |= QItemSelectionModel.SelectionFlag.Deselect
56+
selection_model.select(index, flags)
57+
58+
59+
def _apply_item_expanded(item: QStandardItem, expanded: bool) -> None:
60+
model = item.model()
61+
view = _get_view_from_item(item)
62+
if not (model and isinstance(view, QTreeView)):
63+
item._expanded = expanded
64+
return
65+
66+
index = model.indexFromItem(item)
67+
view.setExpanded(index, expanded)
68+
69+
70+
def _restore_item_state(item: QStandardItem) -> None:
71+
if not hasattr(item, "_state_snapshot"):
72+
return
73+
74+
model = item.model()
75+
view = _get_view_from_item(item)
76+
if not (model and view):
77+
return
78+
79+
snapshot = item._state_snapshot
80+
for it, expanded, _selected in snapshot:
81+
if isinstance(view, QTreeView):
82+
index = model.indexFromItem(it)
83+
view.setExpanded(index, expanded)
84+
for it, _expanded, selected in snapshot:
85+
index = model.indexFromItem(it)
86+
flags = QItemSelectionModel.SelectionFlag.Rows
87+
if selected:
88+
flags |= QItemSelectionModel.SelectionFlag.Select
89+
else:
90+
flags |= QItemSelectionModel.SelectionFlag.Deselect
91+
view.selectionModel().select(index, flags)
92+
93+
del item._state_snapshot
94+
95+
796
@PySideRenderer.register_insert(QStandardItem)
897
def insert(self, el, anchor=None):
998
if hasattr(el, "model_index"):
@@ -19,13 +108,29 @@ def insert(self, el, anchor=None):
19108
break
20109
if index is None:
21110
return
111+
if el.row() >= 0 and el.parent() == self:
112+
self.takeRow(el.row())
22113
self.insertRow(index, el)
23114
else:
24115
self.appendRow(el)
25116

117+
_restore_item_state(el)
118+
119+
if hasattr(el, "_expanded"):
120+
_apply_item_expanded(el, el._expanded)
121+
delattr(el, "_expanded")
122+
123+
if hasattr(el, "_selected"):
124+
_apply_item_selection(el, el._selected)
125+
delattr(el, "_selected")
126+
26127

27128
@PySideRenderer.register_remove(QStandardItem)
28129
def remove(self, el):
130+
snapshot = _snapshot_item_state(el)
131+
if snapshot is not None:
132+
el._state_snapshot = snapshot
133+
29134
if hasattr(el, "model_index"):
30135
# Only support removal of rows for now
31136
row, _column = getattr(el, "model_index")
@@ -48,6 +153,14 @@ def set_attribute(self, attr, value):
48153
model.blockSignals(True)
49154

50155
try:
156+
match attr:
157+
case "selected":
158+
_apply_item_selection(self, value)
159+
return
160+
case "expanded":
161+
_apply_item_expanded(self, value)
162+
return
163+
51164
if attr == "model_index" and model:
52165
index = model.indexFromItem(self)
53166
row, column = value
Lines changed: 135 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,135 @@
1+
<!--
2+
Run with:
3+
uv run collagraph examples/pyside/keyed_tree_demo.cgx
4+
5+
Try this sequence:
6+
1) Expand a couple of parents and select a parent/child in the tree.
7+
2) Click Reverse or Shuffle.
8+
3) Observe that expanded + selection state stays with the moved items.
9+
-->
10+
<widget>
11+
<label text="Keyed TreeView Reorder Demo" />
12+
<label text="Manually expand/select items, then reorder parents." />
13+
14+
<widget :layout="{'type': 'Box', 'direction': 'LeftToRight'}">
15+
<button
16+
text="Reverse Parents"
17+
@clicked="reverse_parents"
18+
/>
19+
<button
20+
text="Shuffle Parents"
21+
@clicked="shuffle_parents"
22+
/>
23+
<button
24+
text="Print UI State"
25+
@clicked="print_ui_state"
26+
/>
27+
</widget>
28+
29+
<treeview object-name="tree">
30+
<itemmodel>
31+
<standarditem
32+
v-for="parent in items"
33+
:key="parent['id']"
34+
:text="parent['text']"
35+
>
36+
<standarditem
37+
v-for="child in parent['children']"
38+
:key="child['id']"
39+
:text="child['text']"
40+
/>
41+
</standarditem>
42+
</itemmodel>
43+
</treeview>
44+
45+
<label :text="status_text()" />
46+
</widget>
47+
48+
<script>
49+
import random
50+
51+
import collagraph as cg
52+
from PySide6 import QtWidgets
53+
54+
55+
class App(cg.Component):
56+
def init(self):
57+
self.state["items"] = [
58+
{
59+
"id": 1,
60+
"text": "Parent 1",
61+
"children": [
62+
{"id": 11, "text": "Child 1.1"},
63+
{"id": 12, "text": "Child 1.2"},
64+
],
65+
},
66+
{
67+
"id": 2,
68+
"text": "Parent 2",
69+
"children": [
70+
{"id": 21, "text": "Child 2.1"},
71+
{"id": 22, "text": "Child 2.2"},
72+
],
73+
},
74+
{
75+
"id": 3,
76+
"text": "Parent 3",
77+
"children": [
78+
{"id": 31, "text": "Child 3.1"},
79+
],
80+
},
81+
{
82+
"id": 4,
83+
"text": "Parent 4",
84+
"children": [
85+
{"id": 41, "text": "Child 4.1"},
86+
{"id": 42, "text": "Child 4.2"},
87+
{"id": 43, "text": "Child 4.3"},
88+
],
89+
},
90+
]
91+
92+
def reverse_parents(self):
93+
self.state["items"] = list(reversed(self.state["items"]))
94+
95+
def shuffle_parents(self):
96+
items = list(self.state["items"])
97+
random.shuffle(items)
98+
self.state["items"] = items
99+
100+
def status_text(self):
101+
parent_count = len(self.state["items"])
102+
child_count = sum(len(parent["children"]) for parent in self.state["items"])
103+
return f"Parents: {parent_count} | Children: {child_count}"
104+
105+
def print_ui_state(self):
106+
tree = self.element.findChild(QtWidgets.QTreeView, "tree")
107+
if not tree:
108+
print("TreeView not found")
109+
return
110+
111+
model = tree.model()
112+
selection_model = tree.selectionModel()
113+
114+
selected = []
115+
expanded = []
116+
117+
for row in range(model.rowCount()):
118+
parent_item = model.item(row)
119+
parent_index = model.indexFromItem(parent_item)
120+
121+
if selection_model.isSelected(parent_index):
122+
selected.append(parent_item.text())
123+
if tree.isExpanded(parent_index):
124+
expanded.append(parent_item.text())
125+
126+
for child_row in range(parent_item.rowCount()):
127+
child_item = parent_item.child(child_row)
128+
child_index = model.indexFromItem(child_item)
129+
if selection_model.isSelected(child_index):
130+
selected.append(child_item.text())
131+
132+
print(f"Order: {[parent['text'] for parent in self.state['items']]}")
133+
print(f"Selected: {selected}")
134+
print(f"Expanded parents: {expanded}")
135+
</script>

0 commit comments

Comments
 (0)