Skip to content

Commit c73d7e3

Browse files
authored
fix(session): use session skill schema for patch merge (#4395)
1 parent 02f5c46 commit c73d7e3

5 files changed

Lines changed: 115 additions & 4 deletions

File tree

openviking/session/compressor_v3.py

Lines changed: 6 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -53,6 +53,10 @@
5353
from openviking.session.memory.utils.json_parser import JsonUtils
5454
from openviking.session.memory.utils.memory_file_utils import MemoryFileUtils
5555
from openviking.session.memory.utils.uri import generate_uri
56+
from openviking.session.skill.session_skill_context_provider import (
57+
SESSION_SKILL_MEMORY_TYPE,
58+
load_skill_extract_registry,
59+
)
5660
from openviking.session.train import (
5761
Case,
5862
ExperienceGradientContext,
@@ -875,7 +879,8 @@ async def _get_session_skill_trainer(
875879
gradient_estimator=_NoopGradientEstimator(),
876880
policy_optimizer=PatchMergePolicyOptimizer(
877881
viking_fs=viking_fs,
878-
memory_type="skills",
882+
memory_type=SESSION_SKILL_MEMORY_TYPE,
883+
memory_registry=load_skill_extract_registry(),
879884
),
880885
policy_updater=SkillPolicyUpdater(
881886
skill_processor=self.skill_processor,

openviking/session/memory/patch_merge_context_provider.py

Lines changed: 5 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -10,6 +10,7 @@
1010

1111
from openviking.server.identity import RequestContext
1212
from openviking.session.memory.dataclass import MemoryFile, MemoryTypeSchema
13+
from openviking.session.memory.memory_type_registry import MemoryTypeRegistry
1314
from openviking.session.memory.session_extract_context_provider import (
1415
SessionExtractContextProvider,
1516
)
@@ -121,9 +122,11 @@ def __init__(
121122
patches: list[PatchMergePatch],
122123
required_file_uris: list[str] | None = None,
123124
output_language: str | None = None,
125+
memory_registry: MemoryTypeRegistry | None = None,
124126
):
125127
super().__init__(messages=[])
126128
self.memory_type = memory_type
129+
self._registry = memory_registry
127130
self.required_file_uris = list(required_file_uris or [])
128131
self.patches = list(patches)
129132
self._output_language = output_language or _resolve_patch_output_language(self.patches)
@@ -153,7 +156,8 @@ def get_tools(self) -> list[str]:
153156

154157
def get_memory_schemas(self, ctx: RequestContext) -> list[MemoryTypeSchema]:
155158
del ctx
156-
schema = self._get_registry().get(self.memory_type)
159+
registry = self._get_registry()
160+
schema = registry.get(self.memory_type)
157161
if schema is None or not schema.enabled:
158162
raise ValueError(f"Memory schema not found or disabled: {self.memory_type}")
159163
return [schema]

openviking/session/train/components/policy_optimizer.py

Lines changed: 4 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -13,6 +13,7 @@
1313
from openviking.session.memory.dataclass import MemoryFile, StoredLink
1414
from openviking.session.memory.extract_loop import ExtractLoop
1515
from openviking.session.memory.memory_isolation_handler import MemoryIsolationHandler
16+
from openviking.session.memory.memory_type_registry import MemoryTypeRegistry
1617
from openviking.session.memory.memory_updater import ExtractContext
1718
from openviking.session.memory.patch_merge_context_provider import (
1819
PatchMergeContextProvider,
@@ -48,6 +49,7 @@ class PatchMergePolicyOptimizer:
4849
viking_fs: Any = None
4950
vlm: Any = None
5051
memory_type: str = "experiences"
52+
memory_registry: MemoryTypeRegistry | None = None
5153

5254
@tracer(
5355
"train.policy_optimizer.patch_merge.plan",
@@ -128,6 +130,7 @@ async def _run_merge_extract_loop(
128130
extract_context = ExtractContext(list(context.messages or []))
129131
provider = PatchMergeContextProvider(
130132
memory_type=self.memory_type,
133+
memory_registry=self.memory_registry,
131134
required_file_uris=_required_file_uris(gradients, policy_set),
132135
patches=[_gradient_to_merge_patch(gradient) for gradient in gradients],
133136
)
@@ -525,7 +528,7 @@ def _name_field_for_memory_type(memory_type: str) -> str:
525528
"""Return the extra_fields key for the policy name in a given memory type."""
526529
if memory_type == "experiences":
527530
return "experience_name"
528-
if memory_type == "skills":
531+
if memory_type in {"skills", "session_skills"}:
529532
return "skill_name"
530533
if memory_type.endswith("s"):
531534
return f"{memory_type[:-1]}_name"

tests/session/memory/test_patch_merge_context_provider.py

Lines changed: 16 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -14,6 +14,7 @@
1414
PatchMergeContextProvider,
1515
PatchMergePatch,
1616
)
17+
from openviking.session.skill.session_skill_context_provider import load_skill_extract_registry
1718

1819

1920
def _memory_file(
@@ -354,6 +355,21 @@ def test_patch_merge_context_provider_get_memory_schema_raises_for_missing_type(
354355
provider.get_memory_schemas(ctx=None)
355356

356357

358+
def test_patch_merge_context_provider_uses_registry_override_for_session_skills():
359+
provider = PatchMergeContextProvider(
360+
memory_type="session_skills",
361+
memory_registry=load_skill_extract_registry(),
362+
required_file_uris=[],
363+
patches=[],
364+
)
365+
366+
schemas = provider.get_memory_schemas(ctx=None)
367+
368+
assert len(schemas) == 1
369+
assert schemas[0].memory_type == "session_skills"
370+
assert schemas[0].enabled is True
371+
372+
357373
def test_patch_merge_context_provider_instruction_mentions_path_field_normalization():
358374
provider = PatchMergeContextProvider(
359375
memory_type="entities",

tests/session/train/test_train_components.py

Lines changed: 84 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -10,6 +10,10 @@
1010

1111
from openviking.session.memory.dataclass import MemoryFile, StoredLink
1212
from openviking.session.memory.utils.memory_file_utils import MemoryFileUtils
13+
from openviking.session.skill.session_skill_context_provider import (
14+
SESSION_SKILL_MEMORY_TYPE,
15+
load_skill_extract_registry,
16+
)
1317
from openviking.session.train import (
1418
ContentHashPolicySnapshotter,
1519
DryRunPolicyUpdater,
@@ -718,7 +722,10 @@ async def run(self):
718722
[],
719723
)
720724

721-
monkeypatch.setattr("openviking.session.train.components.policy_optimizer.ExtractLoop", FakeExtractLoop)
725+
monkeypatch.setattr(
726+
"openviking.session.train.components.policy_optimizer.ExtractLoop",
727+
FakeExtractLoop,
728+
)
722729

723730
plan = await PatchMergePolicyOptimizer(viking_fs=FakeVikingFS({}), vlm=object()).plan(
724731
[gradient],
@@ -1074,3 +1081,79 @@ async def run(self):
10741081
assert captured["constructed"] is True
10751082
assert plan.metadata["patch_gradient_count"] == 1
10761083
assert plan.items[0].after_content == "merged update"
1084+
1085+
1086+
@pytest.mark.asyncio
1087+
async def test_patch_merge_policy_optimizer_uses_session_skill_registry(monkeypatch):
1088+
from openviking.session.memory.dataclass import (
1089+
ResolvedOperation,
1090+
ResolvedOperations,
1091+
)
1092+
1093+
skill_uri = "viking://user/u/skills/code-review/SKILL.md"
1094+
policy_set = ExperienceSet(root_uri="viking://user/u/skills", policies=[])
1095+
gradient = PatchSemanticGradient(
1096+
before_file=None,
1097+
after_file=MemoryFile(
1098+
uri=skill_uri,
1099+
content="Use this skill to review code changes.",
1100+
memory_type=SESSION_SKILL_MEMORY_TYPE,
1101+
extra_fields={
1102+
"memory_type": SESSION_SKILL_MEMORY_TYPE,
1103+
"skill_name": "code-review",
1104+
},
1105+
),
1106+
base_version=None,
1107+
rationale="test",
1108+
links=[],
1109+
confidence=0.9,
1110+
metadata={},
1111+
)
1112+
captured = {}
1113+
1114+
class FakeExtractLoop:
1115+
def __init__(self, **kwargs):
1116+
captured.update(kwargs)
1117+
1118+
async def run(self):
1119+
return (
1120+
ResolvedOperations(
1121+
upsert_operations=[
1122+
ResolvedOperation(
1123+
old_memory_file_content=None,
1124+
memory_fields={
1125+
"skill_name": "code-review",
1126+
"content": "Merged skill content.",
1127+
},
1128+
memory_type=SESSION_SKILL_MEMORY_TYPE,
1129+
uris=[skill_uri],
1130+
)
1131+
],
1132+
delete_file_contents=[],
1133+
errors=[],
1134+
),
1135+
[],
1136+
)
1137+
1138+
monkeypatch.setattr(
1139+
"openviking.session.train.components.policy_optimizer.ExtractLoop",
1140+
FakeExtractLoop,
1141+
)
1142+
1143+
plan = await PatchMergePolicyOptimizer(
1144+
viking_fs=FakeVikingFS({}),
1145+
vlm=object(),
1146+
memory_type=SESSION_SKILL_MEMORY_TYPE,
1147+
memory_registry=load_skill_extract_registry(),
1148+
).plan(
1149+
[gradient],
1150+
policy_set,
1151+
PatchMergePolicyOptimizerContext(request_context=fake_request_context()),
1152+
)
1153+
1154+
assert captured["isolation_handler"].allowed_memory_types == {SESSION_SKILL_MEMORY_TYPE}
1155+
assert len(plan.items) == 1
1156+
assert plan.items[0].memory_type == SESSION_SKILL_MEMORY_TYPE
1157+
assert plan.items[0].target_name == "code-review"
1158+
assert plan.items[0].target_uri == skill_uri
1159+
assert plan.items[0].after_content == "Merged skill content."

0 commit comments

Comments
 (0)