문제 이해
두 개의 이진 탐색 트리(Binary Search Tree, BST)와 하나의 목표 합(sum)이 주어졌을 때, 두 요소의 합이 주어진 값과 같으면서 각 쌍의 요소들이 반드시 서로 다른 BST에 속하도록 하는 모든 쌍을 찾아야 합니다.
예를 들어 합이 12로 주어진 경우를 살펴보겠습니다.

이때 출력은 [(6, 6), (7, 5), (9, 3)]이 됩니다.
접근 방법
이 문제의 핵심 아이디어는 간단합니다. BST를 중위 순회(inorder traversal)하면 항상 오름차순으로 정렬된 리스트를 얻을 수 있습니다. 따라서 두 트리를 각각 중위 순회해 정렬된 두 배열을 만든 뒤 투 포인터(two pointer) 기법을 적용하면, 선형 시간 안에 조건을 만족하는 모든 쌍을 효율적으로 찾을 수 있습니다.
해결 절차는 다음과 같습니다.
- solve() 함수 정의 — 인자로 trav1, trav2, Sum을 받습니다.
- left := 0으로 초기화합니다.
- right := trav2의 크기 − 1로 초기화합니다.
- res := 결과를 담을 새로운 리스트를 생성합니다.
- left < trav1의 크기이고 right ≥ 0인 동안 다음을 반복합니다.
- trav1[left] + trav2[right] == Sum인 경우: (trav1[left], trav2[right])를 res 끝에 추가하고, left는 1 증가, right는 1 감소시킵니다.
- trav1[left] + trav2[right] < Sum인 경우: 합을 키우기 위해 left를 1 증가시킵니다.
- 그 외(합이 Sum보다 큰 경우): 합을 줄이기 위해 right를 1 감소시킵니다.
- 반복이 종료되면 res를 반환합니다.
메인 루틴에서는 다음 작업을 수행합니다.
- trav1, trav2라는 두 개의 새로운 리스트를 생성합니다.
- trav1 := 첫 번째 트리(tree1)의 중위 순회 결과
- trav2 := 두 번째 트리(tree2)의 중위 순회 결과
- solve(trav1, trav2, Sum)을 반환합니다.
구현 예제 (Python)
다음 구현을 통해 동작 과정을 더 자세히 이해할 수 있습니다.
class TreeNode:
def __init__(self, data):
self.data = data
self.left = None
self.right = None
def insert(root, key):
if root == None:
return TreeNode(key)
if root.data > key:
root.left = insert(root.left, key)
else:
root.right = insert(root.right, key)
return root
def storeInorder(ptr, traversal):
if ptr == None:
return
storeInorder(ptr.left, traversal)
traversal.append(ptr.data)
storeInorder(ptr.right, traversal)
def solve(trav1, trav2, Sum):
left = 0
right = len(trav2) - 1
res = []
while left < len(trav1) and right >= 0:
if trav1[left] + trav2[right] == Sum:
res.append((trav1[left], trav2[right]))
left += 1
right -= 1
elif trav1[left] + trav2[right] < Sum:
left += 1
else:
right -= 1
return res
def get_pair_sum(root1, root2, Sum):
trav1 = []
trav2 = []
storeInorder(root1, trav1)
storeInorder(root2, trav2)
return solve(trav1, trav2, Sum)
root1 = None
for element in [9, 11, 4, 7, 2, 6, 15, 14]:
root1 = insert(root1, element)
root2 = None
for element in [6, 19, 3, 2, 4, 5]:
root2 = insert(root2, element)
Sum = 12
print(get_pair_sum(root1, root2, Sum))
입력
[9,11,4,7,2,6,15,14], [6,19,3,2,4,5], 12
출력
[(6, 6), (7, 5), (9, 3)]
복잡도 분석
시간 복잡도: O(n + m) — n과 m은 각각 첫 번째, 두 번째 트리의 노드 수입니다. 중위 순회에 O(n + m), 투 포인터 탐색에 최대 O(n + m)이 소요됩니다.
공간 복잡도: O(n + m) — 두 트리의 중위 순회 결과를 저장하기 위한 추가 공간이 필요합니다.