Compare commits

...

1 Commits

Author SHA1 Message Date
celestial-vault a528373477 make ModelContextTracker a utility function 2025-09-20 22:42:51 -07:00
4 changed files with 57 additions and 55 deletions
@@ -1,40 +0,0 @@
import { getTaskMetadata, saveTaskMetadata } from "@core/storage/disk"
import * as vscode from "vscode"
export class ModelContextTracker {
readonly taskId: string
private context: vscode.ExtensionContext
constructor(context: vscode.ExtensionContext, taskId: string) {
this.context = context
this.taskId = taskId
}
async recordModelUsage(apiProviderId: string, modelId: string, mode: string) {
const metadata = await getTaskMetadata(this.context, this.taskId)
if (!metadata.model_usage) {
metadata.model_usage = []
}
// check to see if the last entry is the same as the new one
const lastEntry = metadata.model_usage[metadata.model_usage.length - 1]
if (
lastEntry &&
lastEntry.model_id === modelId &&
lastEntry.model_provider_id === apiProviderId &&
lastEntry.mode === mode
) {
return
}
metadata.model_usage.push({
ts: Date.now(),
model_id: modelId,
model_provider_id: apiProviderId,
mode: mode,
})
await saveTaskMetadata(this.context, this.taskId, metadata)
}
}
@@ -4,12 +4,11 @@ import { afterEach, beforeEach, describe, it } from "mocha"
import * as sinon from "sinon"
import * as vscode from "vscode"
import type { TaskMetadata } from "./ContextTrackerTypes"
import { ModelContextTracker } from "./ModelContextTracker"
import { recordModelUsage } from "./model-context-tracking"
describe("ModelContextTracker", () => {
describe("recordModelUsage", () => {
let sandbox: sinon.SinonSandbox
let mockContext: vscode.ExtensionContext
let tracker: ModelContextTracker
let taskId: string
let mockTaskMetadata: TaskMetadata
let getTaskMetadataStub: sinon.SinonStub
@@ -28,9 +27,8 @@ describe("ModelContextTracker", () => {
getTaskMetadataStub = sandbox.stub(diskModule, "getTaskMetadata").resolves(mockTaskMetadata)
saveTaskMetadataStub = sandbox.stub(diskModule, "saveTaskMetadata").resolves()
// Create tracker instance
// Set up test data
taskId = "test-task-id"
tracker = new ModelContextTracker(mockContext, taskId)
})
afterEach(() => {
@@ -48,11 +46,12 @@ describe("ModelContextTracker", () => {
const clock = sandbox.useFakeTimers(fakeNow)
try {
// Call the method being tested
await tracker.recordModelUsage(apiProviderId, modelId, mode)
// Call the function being tested
await recordModelUsage(mockContext, taskId, apiProviderId, modelId, mode)
// Verify getTaskMetadata was called with correct parameters
expect(getTaskMetadataStub.calledOnce).to.be.true
expect(getTaskMetadataStub.firstCall.args[0]).to.equal(mockContext)
expect(getTaskMetadataStub.firstCall.args[1]).to.equal(taskId)
// Verify saveTaskMetadata was called with the correct data
@@ -98,8 +97,8 @@ describe("ModelContextTracker", () => {
const clock = sandbox.useFakeTimers(newTimestamp)
try {
// Call the method being tested
await tracker.recordModelUsage(apiProviderId, modelId, mode)
// Call the function being tested
await recordModelUsage(mockContext, taskId, apiProviderId, modelId, mode)
// Verify saveTaskMetadata was called
expect(saveTaskMetadataStub.calledOnce).to.be.true
@@ -158,8 +157,8 @@ describe("ModelContextTracker", () => {
// Reset mock metadata for each iteration to avoid accumulation
mockTaskMetadata.model_usage = []
// Call the method
await tracker.recordModelUsage(provider, model, mode)
// Call the function
await recordModelUsage(mockContext, taskId, provider, model, mode)
// Verify interaction with disk module
expect(getTaskMetadataStub.calledOnce).to.be.true
@@ -0,0 +1,39 @@
import { getTaskMetadata, saveTaskMetadata } from "@core/storage/disk"
import * as vscode from "vscode"
/**
* Records model usage for a task by updating the task metadata
* @param context The VSCode extension context
* @param taskId The ID of the task
* @param apiProviderId The API provider identifier
* @param modelId The model identifier
* @param mode The mode (plan/act)
*/
export async function recordModelUsage(
context: vscode.ExtensionContext,
taskId: string,
apiProviderId: string,
modelId: string,
mode: string,
): Promise<void> {
const metadata = await getTaskMetadata(context, taskId)
if (!metadata.model_usage) {
metadata.model_usage = []
}
// check to see if the last entry is the same as the new one
const lastEntry = metadata.model_usage[metadata.model_usage.length - 1]
if (lastEntry && lastEntry.model_id === modelId && lastEntry.model_provider_id === apiProviderId && lastEntry.mode === mode) {
return
}
metadata.model_usage.push({
ts: Date.now(),
model_id: modelId,
model_provider_id: apiProviderId,
mode: mode,
})
await saveTaskMetadata(context, taskId, metadata)
}
+8 -4
View File
@@ -7,7 +7,6 @@ import { ContextManager } from "@core/context/context-management/ContextManager"
import { checkContextWindowExceededError } from "@core/context/context-management/context-error-handling"
import { getContextWindowInfo } from "@core/context/context-management/context-window-utils"
import { FileContextTracker } from "@core/context/context-tracking/FileContextTracker"
import { ModelContextTracker } from "@core/context/context-tracking/ModelContextTracker"
import {
getGlobalClineRules,
getLocalClineRules,
@@ -68,6 +67,7 @@ import pWaitFor from "p-wait-for"
import * as path from "path"
import { ulid } from "ulid"
import * as vscode from "vscode"
import { recordModelUsage } from "@/core/context/context-tracking/model-context-tracking"
import type { SystemPromptContext } from "@/core/prompts/system-prompt"
import { getSystemPrompt } from "@/core/prompts/system-prompt"
import { HostProvider } from "@/hosts/host-provider"
@@ -119,7 +119,6 @@ export class Task {
// Metadata tracking
private fileContextTracker: FileContextTracker
private modelContextTracker: ModelContextTracker
// Focus Chain
private FocusChainManager?: FocusChainManager
@@ -258,7 +257,6 @@ export class Task {
// Initialize file context tracker
this.fileContextTracker = new FileContextTracker(controller, this.taskId)
this.modelContextTracker = new ModelContextTracker(controller.context, this.taskId)
// Initialize focus chain manager only if enabled
if (this.focusChainSettings.enabled) {
@@ -1689,7 +1687,13 @@ export class Task {
const { model, providerId, customPrompt } = this.getCurrentProviderInfo()
if (providerId && model.id) {
try {
await this.modelContextTracker.recordModelUsage(providerId, model.id, this.mode)
await recordModelUsage(
this.controller.context,
this.taskId,
providerId,
model.id,
this.stateManager.getGlobalSettingsKey("mode"),
)
} catch {}
}