diff --git a/main.ts b/main.ts index 3acfd52..7f08a1b 100644 --- a/main.ts +++ b/main.ts @@ -24,6 +24,7 @@ import { ViewUpdate, WidgetType, } from "@codemirror/view"; +import { Configuration as AzureConfiguration, OpenAIApi as AzureOpenAiApi} from "azure-openai"; import * as cohere from "cohere-ai"; import * as fs from "fs"; import GPT3Tokenizer from "gpt3-tokenizer"; @@ -34,7 +35,7 @@ const untildify = require("untildify") as any; const tokenizer = new GPT3Tokenizer({ type: "codex" }); // TODO depends on model -const PROVIDERS = ["cohere", "textsynth", "ocp", "openai", "openai-chat"]; +const PROVIDERS = ["cohere", "textsynth", "ocp", "openai", "openai-chat", "azure", "azure-chat"]; type Provider = (typeof PROVIDERS)[number]; interface LoomSettings { @@ -42,6 +43,9 @@ interface LoomSettings { cohereApiKey: string; textsynthApiKey: string; + azureApiKey: string; + azureEndpoint: string; + ocpApiKey: string; ocpUrl: string; @@ -68,6 +72,9 @@ const DEFAULT_SETTINGS: LoomSettings = { cohereApiKey: "", textsynthApiKey: "", + azureApiKey: "", + azureEndpoint: "", + ocpApiKey: "", ocpUrl: "", @@ -110,6 +117,7 @@ export default class LoomPlugin extends Plugin { statusBarItem: HTMLElement; openai: OpenAIApi; + azure: AzureOpenAiApi; withFile(callback: (file: TFile) => T): T | null { const file = this.app.workspace.getActiveFile(); @@ -152,6 +160,20 @@ export default class LoomPlugin extends Plugin { cohere.init(this.settings.cohereApiKey); } + // https://github.com/1openwindow/azure-openai-node + setAzureOpenAI() { + const configuration = new AzureConfiguration({ + apiKey: this.settings.azureApiKey, + // add azure info into configuration + azure: { + apiKey: this.settings.azureApiKey, + endpoint: this.settings.azureEndpoint + } + }); + this.azure = new AzureOpenAiApi(configuration); + } + + async onload() { await this.loadSettings(); await this.loadState(); @@ -160,6 +182,7 @@ export default class LoomPlugin extends Plugin { this.setOpenAI(); this.setCohere(); + this.setAzureOpenAI(); this.statusBarItem = this.addStatusBarItem(); this.statusBarItem.setText("Completing..."); @@ -176,6 +199,11 @@ export default class LoomPlugin extends Plugin { !this.settings.openaiApiKey ) return false; + if ( + ["azure", "azure-chat"].includes(this.settings.provider) && + !this.settings.azureApiKey + ) + return false; if (this.settings.provider === "cohere" && !this.settings.cohereApiKey) return false; if ( @@ -252,7 +280,9 @@ export default class LoomPlugin extends Plugin { const loomPanes = this.app.workspace.getLeavesOfType("loom"); try { if (loomPanes.length === 0) - this.app.workspace.getRightLeaf(false).setViewState({ type: "loom" }); + this.app.workspace + .getRightLeaf(false) + .setViewState({ type: "loom" }); else if (focus) this.app.workspace.revealLeaf(loomPanes[0]); } catch (e) { console.error(e); @@ -1106,19 +1136,35 @@ export default class LoomPlugin extends Plugin { top_p: this.settings.topP, }) ).data.choices.map((choice) => choice.text); + } else if (this.settings.provider === "azure-chat") { + rawCompletions = ( + await this.azure.createChatCompletion({ + model: this.settings.model, + messages: [{ role: "assistant", content: prompt }], + max_tokens: this.settings.maxTokens, + n: this.settings.n, + temperature: this.settings.temperature, + top_p: this.settings.topP, + }) + ).data.choices.map((choice) => choice.message?.content); + } else if (this.settings.provider === "azure") { + rawCompletions = ( + await this.azure.createCompletion({ + model: this.settings.model, + prompt, + max_tokens: this.settings.maxTokens, + n: this.settings.n, + temperature: this.settings.temperature, + top_p: this.settings.topP, + }) + ).data.choices.map((choice) => choice.text); } } catch (e) { - if ( - e.response.status === 401 && - ["openai", "openai-chat"].includes(this.settings.provider) - ) + if (e.response.status === 401) new Notice( "OpenAI API key is invalid. Please provide a valid key in the settings." ); - else if ( - e.response.status === 429 && - ["openai", "openai-chat"].includes(this.settings.provider) - ) + else if (e.response.status === 429) new Notice("OpenAI API rate limit exceeded."); else new Notice("Unknown API error: " + e.response.data.error.message); @@ -1220,7 +1266,7 @@ export default class LoomPlugin extends Plugin { completion = completion.replace(/ { const optionEl = providerSelect.createEl("option", { @@ -2365,6 +2414,8 @@ class LoomSettingTab extends PluginSettingTab { dropdown.addOption("ocp", "OpenAI code-davinci-002 proxy"); dropdown.addOption("openai", "OpenAI (Completion)"); dropdown.addOption("openai-chat", "OpenAI (Chat)"); + dropdown.addOption("azure", "Azure (Completion)"); + dropdown.addOption("azure-chat", "Azure (Chat)"); dropdown.setValue(this.plugin.settings.provider); dropdown.onChange(async (value) => { if (PROVIDERS.find((provider) => provider === value)) @@ -2402,6 +2453,18 @@ class LoomSettingTab extends PluginSettingTab { ); apiKeySetting("OpenAI", "openaiApiKey"); + apiKeySetting("Azure", "azureApiKey") + + new Setting(containerEl) + .setName("Azure-Openai Resource Endpoint") + .setDesc("Required if using Azure") + .addText((text) => + text.setValue(this.plugin.settings.azureEndpoint).onChange(async (value) => { + this.plugin.settings.azureEndpoint = value; + await this.plugin.save(); + }) + ); + // TODO: reduce duplication of other settings diff --git a/package-lock.json b/package-lock.json index c9f2d4e..e2ea826 100644 --- a/package-lock.json +++ b/package-lock.json @@ -1,17 +1,18 @@ { "name": "obsidian-loom", - "version": "0.1.0", + "version": "1.6.0", "lockfileVersion": 3, "requires": true, "packages": { "": { "name": "obsidian-loom", - "version": "0.1.0", + "version": "1.6.0", "license": "MIT", "dependencies": { "@codemirror/state": "^6.2.0", "@codemirror/view": "^6.9.3", "@types/lodash": "^4.14.191", + "azure-openai": "^0.9.4", "cohere-ai": "^6.1.0", "gpt3-tokenizer": "^1.1.5", "openai": "^3.2.0", @@ -497,6 +498,15 @@ "follow-redirects": "^1.14.8" } }, + "node_modules/azure-openai": { + "version": "0.9.4", + "resolved": "https://registry.npmjs.org/azure-openai/-/azure-openai-0.9.4.tgz", + "integrity": "sha512-7uii4ZInxzu2zjLg45PdvgOaw3ps18tEAw0Yux9mo8anX4PwnCMSS9xdlKNiNQyyEKPogvAcxH2PIufHXFLx6Q==", + "dependencies": { + "axios": "^0.26.0", + "form-data": "^4.0.0" + } + }, "node_modules/balanced-match": { "version": "1.0.2", "resolved": "https://registry.npmjs.org/balanced-match/-/balanced-match-1.0.2.tgz", diff --git a/package.json b/package.json index ea520cf..ebd6618 100644 --- a/package.json +++ b/package.json @@ -27,6 +27,7 @@ "@codemirror/state": "^6.2.0", "@codemirror/view": "^6.9.3", "@types/lodash": "^4.14.191", + "azure-openai": "^0.9.4", "cohere-ai": "^6.1.0", "gpt3-tokenizer": "^1.1.5", "openai": "^3.2.0",