-
Notifications
You must be signed in to change notification settings - Fork 1
Expand file tree
/
Copy path06_byok_azure_openai.py
More file actions
173 lines (151 loc) · 5.21 KB
/
Copy path06_byok_azure_openai.py
File metadata and controls
173 lines (151 loc) · 5.21 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
#!/usr/bin/env python3
"""BYOK (Bring Your Own Key) — Use Azure OpenAI with the GitHub Copilot SDK.
See the tutorial for learning goals, prerequisites, and usage:
docs/copilot_sdk_tutorial/tutorials/06_byok.md (English)
docs/copilot_sdk_tutorial/tutorials/06_byok.ja.md (日本語)
"""
import argparse
import asyncio
import os
import sys
from azure.identity import DefaultAzureCredential
from _telemetry import add_telemetry_arguments, apply_telemetry_arguments, make_client
from copilot.generated.rpc import PermissionDecisionApproveOnce
from copilot.generated.session_events import (
SessionEventType,
PermissionRequest,
)
from copilot.session import (
PermissionRequestResult,
ProviderConfig,
SystemMessageAppendConfig,
)
def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser(
description="BYOK: Use Azure OpenAI with the GitHub Copilot SDK",
formatter_class=argparse.RawDescriptionHelpFormatter,
epilog=__doc__,
)
parser.add_argument(
"--prompt",
"-p",
default="Briefly explain what BYOK means in the context of AI APIs.",
help="Prompt to send",
)
parser.add_argument(
"--cli-url",
"-c",
default=None,
help=(
"Optional Copilot CLI server URL (e.g. localhost:3000). "
"When omitted, the SDK launches the copilot CLI over stdio."
),
)
parser.add_argument(
"--auth",
choices=["api-key", "entra"],
default="api-key",
help="Authentication method: api-key (default) or entra (Entra ID bearer token)",
)
parser.add_argument(
"--base-url",
default=os.environ.get("BYOK_BASE_URL", ""),
help="Azure OpenAI deployment base URL (overrides BYOK_BASE_URL env var)",
)
parser.add_argument(
"--api-key",
default=os.environ.get("BYOK_API_KEY", ""),
help="Azure OpenAI API key (overrides BYOK_API_KEY env var)",
)
parser.add_argument(
"--model",
default=os.environ.get("BYOK_MODEL", "gpt-4o"),
help="Model/deployment name (overrides BYOK_MODEL env var, default: gpt-4o)",
)
add_telemetry_arguments(parser)
return parser.parse_args()
def _build_entra_bearer_token() -> str:
"""Obtain an Azure Entra ID bearer token via DefaultAzureCredential."""
scope = "https://cognitiveservices.azure.com/.default"
credential = DefaultAzureCredential()
return credential.get_token(scope).token
async def run(
cli_url: str | None, prompt: str, auth: str, base_url: str, api_key: str, model: str
) -> None:
if not base_url:
print(
"Error: --base-url (or BYOK_BASE_URL env var) is required for BYOK mode.",
file=sys.stderr,
)
sys.exit(1)
# ------------------------------------------------------------------
# Build ProviderConfig
# ------------------------------------------------------------------
if auth == "api-key":
if not api_key:
print(
"Error: --api-key (or BYOK_API_KEY env var) is required for api-key auth.",
file=sys.stderr,
)
sys.exit(1)
provider = ProviderConfig(
type="azure",
base_url=base_url,
api_key=api_key,
)
print(f"[Auth] Using API key authentication — model: {model}")
else:
bearer_token = _build_entra_bearer_token()
provider = ProviderConfig(
type="azure",
base_url=base_url,
bearer_token=bearer_token,
)
print(f"[Auth] Using Entra ID bearer token — model: {model}")
# ------------------------------------------------------------------
# Session setup
# ------------------------------------------------------------------
def approve_all(
request: PermissionRequest,
context: dict,
) -> PermissionRequestResult:
return PermissionDecisionApproveOnce()
client = make_client(cli_url)
await client.start()
session = await client.create_session(
on_permission_request=approve_all,
tools=[],
streaming=True,
model=model,
provider=provider,
system_message=SystemMessageAppendConfig(
content="You are a helpful assistant powered by Azure OpenAI."
),
)
print(f"\nYou: {prompt}\nCopilot: ", end="")
def on_event(event) -> None: # noqa: ANN001
if event.type == SessionEventType.ASSISTANT_MESSAGE_DELTA:
print(event.data.delta_content, end="", flush=True)
elif event.type == SessionEventType.SESSION_ERROR:
print(f"\n[Error] {event.data.message}", file=sys.stderr)
session.on(on_event)
await session.send_and_wait(prompt, timeout=300)
print()
def main() -> None:
args = parse_args()
apply_telemetry_arguments(args)
try:
asyncio.run(
run(
cli_url=args.cli_url,
prompt=args.prompt,
auth=args.auth,
base_url=args.base_url,
api_key=args.api_key,
model=args.model,
)
)
except KeyboardInterrupt:
print("\nBye!")
if __name__ == "__main__":
main()